Compare commits

...
80 Commits
Author SHA1 Message Date
Andrew Lee 0bc845d356 vulkan : reuse descriptor sets when bindings are constant (#29280)
* vulkan : reuse descriptor sets when bindings are constant

* vulkan : bump buffer_destroy_count before destroying the buffer
2026-09-29 08:54:55 +02:00
vaibhavdedhiaandAlde Rojas 139997d8e7 chat : fix Muse Glimmer ignoring response_format json_schema with --jinja (#29615)
* chat : fix Muse Glimmer ignoring response_format json_schema with --jinja

Fixes #29613

* chat : accept json fences and clean up

* chat : fix choice parenthesis

---------

Co-authored-by: Alde Rojas <[email protected]>
2026-09-29 08:18:58 +02:00
Adrien Gallouët 76a5bc86d1 common : use fs::path for cache dirs (#29595)
- Avoid useless string conversions on Windows.
- No need for BSD or emscripten special cases.

Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-29 07:24:03 +02:00
Georgi Gerganov 46e17a6352 tests : skip pytest workers when PYTEST_WORKERS=1 (#29610)
Assisted-by: pi:llama.cpp/Qwen3.8-27B
2026-09-29 08:20:54 +03:00
Asahi-PrvandAsahi-Prv fc07d781e6 ci : update the oneAPI toolkit to 2026.1 (#29273)
* ci : update the oneAPI toolkit to 2026.1

oneDNN is removed from Intel Deep Learning Essentials in 2026.0, so
staying on the deep-learning-essentials path would silently lose oneDNN
support when the toolkit version is updated. Switch both the Ubuntu and
Windows CI jobs to the new unified Intel oneAPI Toolkit installer,
which still includes oneDNN (until 2027.0) and keeps the component IDs
unchanged for the Windows install script.

Measured with the same code (b10899) built with oneAPI 2026.1 vs the
2025.3-based release build on Arc B570: prompt processing 1331 vs 434
t/s (3.1x), token generation 50.1 vs 45.3-48.0 t/s.

Assisted-by: GLM (z-ai/glm-5.3-flash)

* docs : update the SYCL backend build requirements for oneAPI 2026.1

With the 2026.0 release the Base toolkit and the HPC toolkit are
combined into the oneAPI Toolkit, and oneDNN is removed from the Deep
Learning Essentials package. Update the install instructions, the
verified release table and the news section accordingly.

Assisted-by: GLM (z-ai/glm-5.3-flash)

* ci : update the release workflow for oneAPI 2026.1 and Level Zero SDK 1.33.1

Align the release package build with the CI build update:
- oneAPI toolkit 2025.3.3 -> 2026.1 (the unified oneAPI Toolkit)
- Level Zero SDK 1.28.2 -> 1.33.1, and the Debian package names
  (level-zero/level-zero-devel -> libze1/libze-dev)
- The Windows DLL copy list for the 2026.1 runtime: sycl9.dll and the
  .6/.3 MKL library versions

Assisted-by: GLM (z-ai/glm-5.3-flash)

* ci : remove the removed .spv fallback files from the Windows DLL copy list

oneAPI 2026.1 no longer ships libsycl-fallback-bfloat16.spv and
libsycl-native-bfloat16.spv (the OpenCL fallback mechanism changed), so
the copy step failed with exit 1.

Assisted-by: GLM (z-ai/glm-5.3-flash)

* devops : update the oneAPI toolkit image in the Intel Dockerfile

Assisted-by: GLM (z-ai/glm-5.3-flash)

---------

Co-authored-by: Asahi-Prv <[email protected]>
2026-09-29 10:19:51 +08:00
Pascal 526c43b8f7 mtmd: fix GCC 15 stringop-overflow in decode_embd_batch (#29607) 2026-09-29 01:21:37 +02:00
Trivikram Reddy 1c4729414d hex-scripts: show trace events smaller than 100nsec in perfetto (#29614) 2026-09-28 15:02:58 -07:00
680a036285 server : support typed content (vision/audio/video) input for /v1/embeddings endpoint (#29556)
* server : support multimodal input for /v1/embeddings (Qwen3-VL-Embedding)

Accept the OpenAI-style wrapped content array format for multimodal
embedding requests. Each {"content": [...]} object is one input that
produces one embedding; text parts are concatenated and image_url parts
are decoded via handle_media then spliced with process_mtmd_prompt.

The legacy formats (plain string, token arrays, mixed arrays, and the
{prompt_string, multimodal_data} object) continue to work unchanged via
tokenize_input_prompts. Bare content arrays (the unwrapped shape) are
rejected with a migration message.

Also disables KV prefix reuse for stateless embedding/rerank tasks so
that repeated inputs do not incorrectly share cached KV across requests.

Assisted-by: Opencode Qwen3.8 27B

* clean up comments and docs

* refactor

* add tests

* support video and audio inp

---------

Co-authored-by: timothywang21 <[email protected]>
Co-authored-by: Xuan Son Nguyen <[email protected]>
2026-09-28 21:40:38 +02:00
Linux User 66e665c427 vulkan: include functional header (#29597)
Fixes the compile error: `no template named 'function' in namespace 'std'`
2026-09-28 21:31:42 +02:00
Pascal 57b557cb95 models: pad on the left with ggml_pad_ext (#29567)
* models: pad on the left with ggml_pad_ext

The Parakeet, LFM2-Audio, Granite Speech and Gemma 4 audio encoders
build a left padding as a right pad followed by a roll, and DFlash2
concatenates a zero filled block in front of the previous tokens.
ggml_pad_ext does both in one node now that every backend supports a
left padding. The Gemma 4 audio embeddings are bit identical.

* models: skip the DFlash2 taps that only read padding

A tap at or past block_size shifts every row out of the block, so its
term is zero. The loop runs min(kernel_size, block_size) taps.
2026-09-28 20:56:15 +02:00
Ravi Panchumarthy 14ebbd5f2f ggml-openvino: mark unaligned batch-stride views unsupported (#29603) 2026-09-28 21:54:01 +03:00
Xuan-Son Nguyen f1ea206218 batch: migrate speculative, mtmd and server to batch_ext (#29385)
* adapt common

* add common_batch

* wip

* wip: spec

* cont

* common_speculative_process

* server_batch to use common_batch

* rm some stale calls

Assisted-by: Claude Fable 5.1

* migrate mtmd

* handle imrope, handle return val of add()/add_embd()

* add spec zeros vector

* add warning on zero fill path
2026-09-28 19:52:45 +02:00
Adrien Gallouët 6c7a87f7e5 common : fix HF cache paths on Windows (#29475)
Supersedes #29158

Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-28 16:25:24 +02:00
jbooth f00a64c147 webgpu: Handle unaligned writes in ggml_backend_webgpu_buffer_set_tensor (#29471)
* Fix:  Handle unaligned writes in ggml_backend_webgpu_buffer_set_tensor

* Clang formatting
2026-09-28 16:38:10 +03:00
Georgi Gerganov d77dd0806d tests : refactor test-recurrent-state-rollback (#29426)
* tests : use llama_context_ptr in test-recurrent-state-rollback

Replace raw llama_context pointers with llama_context_ptr and drop the
manual llama_free calls and cleanup lambda.

Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-RL

* tests : run test-recurrent-state-rollback over all dummy models

Add a --models DIR mode that mirrors test-save-load-state: iterate every
dummy model, report PASS/FAIL/SKIP in a table and fail only when a model
fails. Register a single ctest entry with ARGS --models instead of the four
per-model registrations.

Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-RL

* cont : fix typo

* metal : allow fusing 0-element nodes to keep graph packing shape-independent

The fusion packing in ggml_metal_fusion_max excluded 0-element tensors and
the topk_moe/moe_reduce checks rejected n_tokens == 0, so graphs decoding
batches with no outputs packed differently from the worst-case reserved
graph. The Metal optimizer then reordered the nodes differently and
ggml_gallocr_needs_realloc failed on the layout mismatch, forcing an
unexpected graph re-reserve (caught by GGML_SCHED_DEBUG_REALLOC).

Treat empty tensors like their non-empty counterparts: match them in the
pattern sequence and only reject genuinely malformed shapes. Fused kernels
dispatch zero threadgroups for empty graphs, which is a legal no-op.

Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-RL

* tests : run test_multi_seq_split_replay as a separate test

test_multi_seq_split_replay was invoked at the end of test_rollback,
so its result was folded into the rollback status and it only ran when
the rollback part passed.

Give it its own test_status return, run both tests independently over
both cache fills via a shared run_tests helper, and report them as
separate rollback / split replay columns in the --models table with
per-test summaries. The exit code fails when either test fails.

Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-RL

* tests : loosen the split replay nmse bound to 1e-4

test-generate-models seeds its weights from std::random_device, and some
generated lfm2 models drift up to ~1.7e-5 nmse on the split replay due to
rounding noise, tripping the previous 1e-5 bound. Raise the bound to 1e-4
so the random generations stop flaking.

Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-RL

* tests : reuse run_tests_for_model in single-model mode

The single-model path duplicated the model init and the non-recurrent
check from run_tests_for_model; route it through the shared helper
instead. Model load failures now return FAIL rather than SKIP so that
--model with a broken file still exits non-zero, and the helper loads
with model_only like the --models loop does since the tests create
their own contexts.

Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-RL
2026-09-28 16:36:38 +03:00
SXX 6f767fe960 ggml-cpu: enable tiled flash attention for non-vector-multiple head dims on x86 (#29423)
* ggml-cpu: enable tiled flash attention for non-vector-multiple head dims on x86

* add AVX2 support for masked loading and storing in simd_gemm_ukernel_tail

* ggml-cpu: fix FA softcap handling for padded KV tiles
2026-09-28 16:23:31 +03:00
uvos f916130d00 ci : ignore more vgpr spills in > 256 DQK fattn kernels (#29571) 2026-09-28 14:52:07 +02:00
François-Xavier Gsell 03a667aa30 vulkan: fuse qwen4exp's SCALE -> SIGMOID -> SCALE -> hc_post chain (#29520) 2026-09-28 14:20:14 +02:00
uvos c2a9e16068 HIP: fix template skip for DKQ > 256 mfma kernels (#29559) 2026-09-28 13:47:29 +02:00
Pascal 4364bf7232 metal: support left and circular padding in GGML_OP_PAD (#29561)
* metal: support left and circular padding in GGML_OP_PAD

Align Metal with CPU, CUDA and Vulkan: shift the source coordinates by
the left paddings, wrap them around with the same wrap_around when
circular, and read the source through nb00, which also fixes a right
padding of a permuted source. A test case covers it.

Drop the f32_4 kernel: its selection is disabled as slower, and it
fails two pad cases once enabled.

* metal: use a function constant for the circular pad variant

Address review from ggerganov: replace the bool template with FC_PAD,
as FC_upscale_aa does, so the pad kernel is compiled once and
specialized per pipeline.
2026-09-28 12:26:50 +02:00
Sihan YuandGeorgi Gerganov ed7ac35e1e context : do not re-reserve the scheduler when toggling causal_attn (#28751)
* context : do not re-reserve the scheduler when toggling causal_attn

`llama_context::set_causal_attn()` marks the scheduler to do a full re-reserve on every change of the flag. For vision inputs, this flag is flipped twice around each non-causal image chunk for Gemma models, resulting in two expensive `sched_reserve()` passes per image. This is especially slow for multi-image or video inputs.

The cost of a re-reserve scales with context and ubatch configurations, so larger settings pay more per image (see table below).

The re-reserve is unnecessary in this case because `causal_attn` only changes the values written to KQ mask, not tensor shapes or any other buffer sizes.

Note: `causal_attn` is a graph reuse key (`llm_graph_params` via `cparams`), so a new graph is built regardless of `sched_need_reserve`, so this doesn't change the graph rebuilding behaviour.

llama-server with gemma-4-26B-A4B Q4_0 + BF16 mmproj, 130-token images,
cache_prompt=false, prompt_ms median of 3 (before -> after):

| images | config | H200 before -> after | RTX 4090 before -> after |
|-|-|-|-|
| 1 | `-c 8192 -ub 512` | 134 -> 105 ms (1.27×) | 201 -> 119 ms (1.69×) |
| 24 | `-c 8192 -ub 512` | 2278 -> 1562 ms (1.46×) | 3559 -> 1748 ms (2.04×) |
| 24 | `-c 32768 -ub 2048` | 5379 -> 1584 ms (3.40×) | 13377 -> 1759 ms (7.61×) |

Generated output remains identical before and after.

* qwen4exp : make the indexer bias shape independent of causal_attn

The block/cell bias path was selected on cparams.causal_attn, so the
causal and non-causal graphs differed in tensor shapes and ops. With the
re-reserve removed (previous commit), a runtime flip resulted in
reallocating the compute buffers, which would fail under
GGML_SCHED_NO_REALLOC.

This commit selects the block path from the mask shape only, independent
of causal_attn. causal_attn is instead passed to set_input_qsa.
causal_attn is fixed per graph as it's part of the reuse key. Causal
values are unchanged. Non-causal values now follow the reference rule,
where every visible block competes on score and only unpooled cells are
always selected.

* context : state the causal_attn shape rule in the comment

* cont : add TODOs

---------

Co-authored-by: Georgi Gerganov <[email protected]>
2026-09-28 11:58:51 +03:00
Sarah Wu 0c6a6a7ce5 Enables Windows ARM64 build with MSVC cl.exe (#28362)
* can reproduce the issue vlad sees

* fix fma issue

* drop volatile

* fix volatile runtime task

* add arm flag if needed

* fix hsum compile error

* fix syntax in quants

* strengthen sve probing

* make the syntax fixes one liners

* remove debug code

* formatting

* remove macro for float

* drive down gcc instruction count

* support armec

* fix CI comments

address CI comments

fix cross compile issue

remove warning

fix style and fix fma probing

fix style

* add documentation

* update documentation
2026-09-28 10:07:27 +02:00
Georgi Gerganov 81ef10ea58 tests : fix ggml init (#29554)
* tests : init ggml for test-recurrent-state-rollback

* cont : same for test-save-load-state

* cont : add to test-state-restore-fragmented + add TODOs
2026-09-28 10:24:21 +03:00
Pascal 5262471615 vulkan: fix wrong results when a mul_mat reads a slice of a larger cache (#28956)
* vulkan: read the batch stride of an in place src0 from nb[2]

A dim01 contiguous tensor can still be a view whose batches are
strided by more than ne[1] rows, the first rows of a KV cache for
example. Both the mat-vec and the matrix paths read such a tensor in
place but passed ne00*ne01 as the batch stride, so every head past
the first read the wrong rows. The same applies to src1. The stride
now comes from nb[2] whenever the tensor is used in place; the value
is unchanged for a contiguous tensor.

test-backend-ops gets an m_v parameter on test_mul_mat, the number of
rows of a in memory, and two cases at the shapes of a decoder self
attention over a cache.

* vulkan: size the in place A and B ranges by their strided extent

The matrix path bound src0 and src1 to the shader with a range of
elements times type size, which ends before the batches of a strided
view. Pipelines with bounded access read zero past that range, so the
same view that the mat-vec path already handles gave wrong results
on Intel and on NVIDIA without coopmat2. The range now comes from
ggml_nbytes when the tensor is read in place.

* vulkan: address review from jeffbolznv

Bind the in place A and B of the matrix path with ggml_vk_subbuffer,
which spans to the end of the buffer, so a strided view is in range
without computing its extent.

mul_mat_id reads the batch stride of an in place src0 and src1 with
the same helper as mul_mat. test_mul_mat_id gets an m_v parameter,
the number of rows of as in memory, and a case whose experts are
strided by more rows than it uses.

* vulkan: read the batch stride of an in place src0 in mul_mat_vec_id

The single token path of mul_mat_id passed ne00*ne01 as the batch
stride of A, so a strided expert view read the wrong rows. The stride
now comes from ggml_vk_batch_stride like the other three paths, and
src1 follows the same rule.

test_mul_mat_id gets a single token case over the strided view.

* vulkan: address review from jeffbolznv

The batch stride of an in place tensor is taken from nb[2] as
nb[2] / type_size * block_size, which holds when nb[2] is padded and
not a multiple of nb[1]. A test_mul_mat case with a padded batch stride
covers it.

* vulkan: keep the A and B ranges exact in mul_mm

The quantized A loads of mul_mm carry no row bound and rely on the
descriptor range to read zeros past the last row of a partial tile.
Binding A and B up to the end of the buffer let those tiles read the
leftovers of a previous node and hung the NVFP4 mul_mm on NVIDIA
without coopmat2. The range is the strided extent of a tensor read in
place and the staged size otherwise.
2026-09-28 08:36:38 +02:00
Tim Wangandtimothywang21 4da6337767 server : allow RANK pooling batch splitting for causal LLM rerankers (ie. Qwen3 and Qwen3-VL) (#28876)
* server : allow splitting RANK pooling for causal LLM rerankers

Rerank models fall into two categories: bidirectional cross-encoders
(BERT, etc.) that require all tokens in a single physical batch, and
causal LLMs repurposed as rerankers (Qwen3, Qwen3-VL) that can use
chunked prefill like any other decoder.

Previously the server rejected all RANK-pooling inputs larger than
n_ubatch, and the graph builder hardcoded QWEN3/QWEN3VL arch checks to
determine last-token pooling. This broke long-document and multimodal
reranking for causal models.

Fix: expose llama_get_causal_attn(ctx) so the server can check the
effective runtime attention type (reflecting any --attention override
or set_causal_attn call). Also expose llama_model_is_causal(model)
for querying the static architectural property from GGUF metadata.

can_split() now permits chunked prefill for RANK pooling when the
context is causal. The graph builder's inline arch check is replaced
with the same cparams.causal_attn predicate, removing the duplication.

Assisted-by: Opencode/Qwen3.8-27B

* remove unused llama_model_is_causal, fix whitespace

Assisted-by: opencode

---------

Co-authored-by: timothywang21 <[email protected]>
2026-09-27 23:28:10 +02:00
Georgi Gerganov a97cce86a8 common : avoid side effects around params parsing (#29537)
- register --rpc unconditionally and call llama_supports_rpc() only from its handler
- print server "initialization ..." log after args are parsed

Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-RL
2026-09-27 20:18:56 +03:00
Adrien Gallouët 136887b665 common : make string_split<T> throw on invalid values (#29518)
Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-27 18:04:11 +02:00
Toki Nasin 9adc7f420c convert : export YaRN scaling parameters for PLaMo-3 (#29528)
Recent PLaMo-3 models use YaRN, while some earlier PLaMo-3 models do not.
The recent PLaMo-3 store their YaRN settings as flat config keys
(rope_scaling_factor, initial_context_length) and build the dict at runtime
in Plamo3Config.rope_parameters. The current converter misses these settings
and writes plain RoPE metadata to GGUF. Mirror the runtime settings into
rope_parameters so the corresponding rope.scaling.* is written to GGUF.
2026-09-27 13:45:47 +02:00
Sigbjørn Skjæret 6fd50a4094 ci : bump ty to 0.0.84 (#29529)
* bump ty to 0.0.84

* fix assertion bug caught by ty
2026-09-27 13:43:08 +02:00
Sigbjørn Skjæret 33c923db1b jinja : add support for dict builtin (#29477)
* add support for dict builtin

* add tests
2026-09-27 13:41:50 +02:00
lhez c9064dded7 opencl: refine bin kernel loading condition (#29503) 2026-09-27 13:08:39 +03:00
bri-prism c829670992 sycl: FWHT kernels for block widths above 512 (#29243)
The SYCL FWHT covers 64 to 512 via the standard butterfly network, plus
384/640/768/1280 via the Kronecker/Paley construction added separately in
Hadamard hint can produce (1024, 2048, 4096, 8192); those still fall through
to the default case and run as a dense GEMM against the materialized
rotation tensor, correct but O(n^2) instead of O(n log n).

fwht_kernel_wide runs one row per work-group instead of per sub-group, so
each work-item keeps N/NT values rather than N/WARP_SIZE. Butterflies below
the sub-group width still shuffle; those up to the work-group width go
through work-group local memory; the rest stay in registers. Same butterfly
and sign convention as the existing narrow kernel.

ggml's SYCL backend registration (dpct::dev_mgr) unconditionally requires a
GPU-labeled platform to exist and throws before any op-level test can run,
so test-backend-ops could not be exercised on this box (a GPU-less pod) even
via the CPU device. Verified instead with a standalone harness: the same
kernel body run through a real SYCL CPU device (Intel oneAPI DPC++ 2026.1,
OpenCL CPU backend), checked against an independent recursive-doubling
Hadamard reference, cross-validated by first running the existing unmodified
narrow kernel through the identical harness and confirming it passes (rules
out a reference-convention bug before trusting a pass on the new code).
Random-input results for all four widths, single- and multi-row:

  N=1024 NT=256 rows=1  max_abs_err=1.7e-07  max_rel_err=4.9e-04  PASS
  N=2048 NT=256 rows=1  max_abs_err=1.9e-07  max_rel_err=2.0e-04  PASS
  N=4096 NT=256 rows=1  max_abs_err=2.0e-07  max_rel_err=1.4e-04  PASS
  N=8192 NT=256 rows=1  max_abs_err=2.5e-07  max_rel_err=3.8e-03  PASS
  N=1024 NT=256 rows=7  max_abs_err=2.4e-07  max_rel_err=1.0e-03  PASS
  N=2048 NT=256 rows=5  max_abs_err=3.0e-07  max_rel_err=9.4e-04  PASS
  N=4096 NT=256 rows=3  max_abs_err=2.7e-07  max_rel_err=1.7e-03  PASS
  N=8192 NT=256 rows=2  max_abs_err=2.5e-07  max_rel_err=1.9e-03  PASS

This covers the kernel algorithm itself; it does not exercise the ggml
dispatch/supports_op integration end to end, which needs a real GPU (or a
SYCL GPU plugin) to get past backend registration. test-backend-ops build
is verified: fwht.cpp recompiles with zero warnings as part of ggml-sycl.
2026-09-27 13:08:19 +03:00
Animesh 36d7b08340 CUDA: tune fp16 tile FlashAttention configs for head sizes 40-112 (#26289) 2026-09-27 13:07:37 +03:00
uvos 2ebd9ae621 HIP: Enable fattn-mma kernel on cdna for dkq > 256 for large batch sizes (#28907)
* HIP: Enable fattn-mma kernel on cdna for dkq > 256 for large batch sizes

* CI: hip-quality-check: ignore spills for very large mfma mma kernels
2026-09-27 13:06:56 +03:00
Ruben Ortlam cea74625fa vulkan: fix argsort kernel selection for Adreno (#29469) 2026-09-27 13:06:06 +03:00
Adrien Gallouët da6c28eb13 common : throw instead of abort on grammar without llguidance (#29516)
Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-27 13:05:50 +03:00
Aman Gupta d7fb90e8e2 RPC: use RDMA completion channel to not spin (#29440)
* RPC: use RDMA completion queue to not spin

* add TODO for apple RDMA
2026-09-27 17:28:32 +08:00
Georgi Gerganov 7fb2b082ce ci : enable GGML_SCHED_DEBUG_REALLOC=1 for ctest workflows (#29514)
* ci : enable GGML_SCHED_DEBUG_REALLOC=1 for ctest workflows

* cont : metal paravirtual device is not compatible
2026-09-27 12:16:47 +03:00
Adrien Gallouët 187664b537 llama-bench : fix OOB access of hf_file (#29515)
Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-27 10:24:54 +02:00
Georgi Gerganov 85ca3b52c3 hrm : fix layer placement of z_l_init weight (#29512) 2026-09-27 10:16:23 +03:00
kurquharandMax Krasnyansky 7ac59a6e3a hexagon: support tiled Q4_0 and Q8_0 GET_ROWS (#29511)
* hexagon: support tiled Q4_0 and Q8_0 GET_ROWS

* hex-get-rows: fix macros

* hex-get-rows: use tiled HVX dequantization

Assisted-by: OpenCode

* hex-get-rows: fix register spills and clean up checks for unsupported ops

* hex-get-rows: improve dma pipeline

* hex-get-rows: improve/simplify kernel selection logic

* hex-build: reenable vectorizer, didnt notice the regression earlier in the sampler update

---------

Co-authored-by: Max Krasnyansky <[email protected]>
2026-09-26 23:02:37 -07:00
Max Krasnyansky 2b129ccfa0 hexagon: support for backend sampler (#29502)
* hex-topk: trying to improve/cleanup the pipeline

* hex-sampling: add STEP op

* hex-sampler: add SUM op

* hex-sampler: update CPY to support sampling cases

* hex-binary: add support for chunking to handle large logits

* hex-argmax: super basic version of ARGMAX

* hex-binary: support for scalars in extended buffers

* hex-binary: fix wrong indexing for dim 1 broadcasts across dim 2 slices

* hex-argsort: fix missing header

* hex-sampler: cleanup dma usage in the sampler related ops, and binary

* hex-build: disable autovectorizer, it is better to use explicit hints for critical loops

* hex-binary: fix perf regression due to is_1d fallback

* hex-ops: update supported ops
2026-09-26 20:29:36 -07:00
Anav Prasad 95887577ab cuda: support Nemotron 3 Puzzle state size 96 for ssm scan (#28717) 2026-09-26 22:05:52 +02:00
R0CKSTAR 694ec23548 musa: build the docker images from the PH1 MUSA SDK image (#29481) 2026-09-26 21:50:27 +02:00
bri-prism 6f856c7099 cuda: add F16 input to the FWHT (#29096)
* cuda: add F16 input to the FWHT

The CUDA FWHT accepts F32 input only. This makes the source type a template
parameter, so the kernel reads an F16 source directly instead of requiring a
converted copy. The F32 path is unchanged.

supports_op accepts an F16 src1 against an F32 src0 for the Hadamard hint.
Every other F16 src1 against a non-F16 src0 is still refused.

ggml_cuda_op_mul_mat_use_fwht is the single predicate both supports_op and
the dispatch call now share, checking contiguity and same-shape(src1, dst)
in addition to the type/hint conditions above. Without a shared predicate,
supports_op could admit an op that ggml_cuda_op_fwht then rejects only after
the unconditional same-shape assert has already fired; that gap predates
this change (it applies to the existing F32 path too) but this PR is what
touches supports_op, so it closes it here.

test-backend-ops on an A10 (lambdalabs): MUL_MAT 1297/1297, including all
24 Hadamard cases (18 existing F32, 6 new F16).

* cuda: use ggml_cuda_cast in the FWHT load, drop the comment
2026-09-26 21:36:09 +02:00
Adrien Gallouët fcb3074f2b server : fix wake_fd warning on Windows (#29479)
Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-26 21:25:25 +02:00
Gaurav Garg 2145525a40 Revert "Change max context length for auto-fitting with unified KV (#28849)" (#29437)
This reverts commit b04d4e567c.
2026-09-26 07:59:15 -07:00
Sigbjørn Skjæret 81bc6b83f8 jinja : implement sameas test (#29448)
* implement sameas test

* add tests
2026-09-26 10:14:58 +02:00
Sigbjørn Skjæret 86a24a182b jinja : fix compile error (#29468) 2026-09-26 09:43:47 +02:00
Chipmunk 08618ff8e7 llama : fix K/V and recurrent state cleanup after failed restores (#27530)
* llama : add discard for deferred state writes

* llama : add tensor zeroing helper for backends without tensor memset

* llama : clear K/V data after failed sequence restore

* llama : clear recurrent state data after failed sequence restore

* llama : simplify discard and restore cleanup

* llama : report error when abnormal cell count is found in state_read_meta

* llama : clear attention state on hybrid restore failure

* tests : cover failed state restore cleanup

* llama : clear MLA state on dsa restore failure

* tests : update test for rebased test suite

* llama : clarify comment in llama_memory_recurrent::state_read
2026-09-26 10:23:03 +03:00
Sigbjørn Skjæret a1de614ba3 jinja : support noncall test statements with arg (#29443)
* support noncall test statements with arg

* add tests
2026-09-26 08:56:48 +02:00
Vladislavandplotnikov.v10 965f89794f polished Readme and llama-bench (#28968)
Co-authored-by: plotnikov.v10 <[email protected]>
2026-09-26 14:21:21 +08:00
jboothandGeorgi Gerganov d834d44e64 ggml-cpu: tiled mul_mat for k-quants (#27851)
* Added tiled mul_mat.

For each mul_mat_one_chunk, quants are unpacked into (max) 256x256 tiles of int8,
one routine per quent.  Then microkernel computes 16x16 tiles before writing out
256x256 float reults to main memory.

Tests/benches in tests/test-tiled-mulmat.cpp.  3-6x speed improvement
for large matmul, break even at 4096x64 * 64x4096, 80% performance (net
loss) for GEMV.  Error rates trivial (order of 1-e04 max, 1-e05 rmse).

* Fixes for ARM/windows builds

* more windows fixes, ggml-cpu.h isn't visible in MSVC for some reason

* unified iqp + tiled on the Q5_K, IQ4_XS set for benchmarking, updated benchmark

* Fixed accidental removal of llama_build_and_test(test-backend-ops.cpp)

* First integration of iqp code

Co-authored-by Bartowski <[email protected]>

* Cleaning up declaration of iq unpacking helpers to align with the bit unpackers

* Removed iqp path

* Fix cross-platform warnings

* Disabling benchmarks unless explicitly enabled

* Fix backend_init for DLL-based builds, add self and bartowski to CODEOWNERS for tiled

* Put benchmarks behind a flag

* kernel fix for AVX2, iq quants

* Fix for asan, leaking memory in test-tiled-mulmat and avoid stack use after return

* guarding env flags with std::call_once

* Simplified repacking for VNNI to a single call per macrotile

* No threadlocals anymore, aligned wdata access

* Doing aligned reads since we ensure alignment with padding in wdata

* Eliminated per-thread gather of Q8_K rows in mul_mat_id, we now gather/repack in a single pass.  Repack method now takes pointer array to support both dense/normal and mmid paths.  Interface with ggml-cpu.c simplified as a result

* Unified/simplified dispatch and support checks.  Put details on wdata needed inside the kernel.h body, simplified interactions with ggml-cpu.c.

* Cleanup includes and whitespace, update src1_repack to return false if we don't need a special repack, so the common case is handled by driver

* Better detection of win32 and additional whitespace fixes

* Gating fuzz tests behind a parameter and some extra prints to try and fix slow CI hosts

* Optimized AVX2 kernel

* Changed interleave format and added ability to interleave in-place after dequant

* Repacks now happen in-place, 16x64 microtiles are independent of each other

* Only repack rows in groups of 16 as they're needed.  Save work in low n_rows cases and optimize L1 usage in other cases

* Use long panels for memory-bound regime (M <= 16), reintroduce IQP path for benchmarks

* Fix unused warnings and cleanup.  Improved IQ dequantization speed.

* Removed separate process benchmarks

* Revert "Removed separate process benchmarks"

This reverts commit 0688cf43d5.

* AVX2 optimizations and guards for tests on windows

* Removed temp perf harness

* Remove perf-mulmat from build

* Removed IQP path, simplified tests to not use sub processes

* Cleaning up alignment of wdata

* Whitespace fixes and aligning L2 workspace to clean 512kb boundaries

* Update ggml/src/ggml-cpu/tiled/tiled-kernel.cpp

Co-authored-by: Georgi Gerganov <[email protected]>

* Cleanup merge-duplicated declaration of test-backend-ops target

* Undo accidental line deletion in ggml.c

---------

Co-authored-by: Georgi Gerganov <[email protected]>
2026-09-26 08:38:57 +03:00
shaofeiqi 9f70b2cecd opencl: add A8 Q8_0 non-MoE dp4a binary kernel (#29439) 2026-09-25 20:41:47 -07:00
Trivikram Reddy 4e7481175c hexagon: find software divide calls using binary inspection tool (#29449)
* hex-scripts: fix table alignment

* hex-scripts: find sw div calls using binary inspection tool
2026-09-25 19:43:38 -07:00
Alessandro de Oliveira Faria (A.K.A.CABELO) 171e8846b4 vendor : update cpp-httplib to 0.58.0 (#29407) 2026-09-26 01:55:18 +02:00
Adrien Gallouët 4b1a27fa0e common,rpc : simplify fs_create_directory_with_parents() (#29432)
The original function was broken on Windows for some unicode paths

Paths without a trailing separator now create the last directory too,
matching the function name. All current callers already include a
trailing separator, so this change does not affect them.

Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-25 20:33:37 +02:00
Yuri Khrustalev fcc891545b mtmd: fix mel preprocessor in LFM2 audio (#29403)
which resulted in different greedy transcripts for 4.5% of English and 6.5% of Japanese
test utterances. In Japanese, some differences changed entire words.

This change:

* uses `log(x + 2^-24)` instead of clamping to the log floor
* uses a symmetric Hann window, equivalent to `torch.hann_window(periodic=False)`
* adds the normalization epsilon to the standard deviation instead of inside the square root

Only the `lfm2a` preprocessor opts into these behaviors. Other audio preprocessors are unchanged.

Tested on top of 84e76d8 using `llama-server` with CUDA and `temperature=0`, compared against
http://github.com/Liquid4All/liquid-audio fp32.

Test set:

* 200 LibriSpeech `test-clean` utterances (EN)
* 200 Common Voice `ja` test utterances (JP)
* identical 16 kHz audio passed to both implementations

| Greedy transcript identical to `liquid-audio` | Without fix |    With fix |
| --------------------------------------------- | ----------: | ----------: |
| EN F16                                        |     191/200 | 200/200 |
| JP F32                                        |     187/200 | 200/200 |
| JP F16                                        |     187/200 | 199/200 |

The remaining JP F16 difference is a comma and matches the reference implementation's own bf16
output.

Mel relative L2 error versus `liquid-audio`:

* EN: 3.2% -> ~2e-6 median
* JP: 3.9% -> ~2e-6 median
2026-09-25 18:52:44 +02:00
shaofeiqiandLi He a25c9865fe opencl: add bin kernel kernel_gemm_noshuffle_q5_k_f32_32b_trans_ila_a8_bin, kernel_gemm_noshuffle_q5_k_q8_1_dp4a_ila_a8_bin (#29401)
* opencl: add A8 Q5_K non-MoE non dp4a + dp4a binary kernel

* opencl: fix s transpose - s only transposed for bin kernels

---------

Co-authored-by: Li He <[email protected]>
2026-09-25 07:35:18 -07:00
sliu39 e85e15cf6d Fixing the vulkan build issue of legacy GLSLC version that has no cooperativeMatrix API support (https://github.com/ggml-org/llama.cpp/issues/29373) (#29409)
* vulkan : fix build issue of legacy glslc version by adding GGML_VULKAN_COOPMAT_GLSLC_SUPPORT macro check for Intel FA shader compiling

* vulkan : add preprocess condition to filter out unsupported FA 2 phases kernels before creation.

* vulkan : move lock_guard for Intel FA shader pointer creation under CM1 compiling preprocessor
2026-09-25 14:33:49 +02:00
Sigbjørn Skjæret b248f4a3c1 gguf-py : ByteLevel processing defaults bos/eos to False (#29422) 2026-09-25 13:59:47 +02:00
Sigbjørn Skjæret d81aef1994 gguf-py : TemplateProcessing has final word on add_special_token (#29417)
* templateprocessing must win over tokenizer config

* remove obsolete override
2026-09-25 11:55:38 +02:00
Adrien Gallouët 27b20ba8b1 common : extract shared unicode path/string helpers (#29415)
Signed-off-by: Adrien Gallouët <[email protected]>
2026-09-25 11:40:41 +02:00
bri-prismandGeorgi Gerganov e351231c4f metal: FWHT kernels for block widths above 512 (#29095)
* metal: FWHT kernels for block widths above 512

The Metal FWHT covers widths 64 to 512, one row per simdgroup with N/32 values
per lane. Wider blocks need more registers per lane than that layout allows.

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

Widths 64 to 512 keep the simdgroup kernel. 1024 through 8192 use the new one,
for both F32 and F16 sources.

The wide kernels allocate float[N] of threadgroup memory, 32 KB at 8192, so the
size check takes the device limit and reports those widths as unsupported where
they would not fit. Without that a device with less threadgroup memory would
accept the op and then abort on a nil pipeline.

test-backend-ops on M5 Pro: MUL_MAT_HADAMARD 26/26, MUL_MAT 1265/1265.

* cont : add TODOs

---------

Co-authored-by: Georgi Gerganov <[email protected]>
2026-09-25 12:15:33 +03:00
Georgi Gerganov 5a75f14c0f metal : split fa kernels into per-dtype libraries (#29329)
* metal : split fa kernels into per-dtype libraries

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

* cont : minor fix comment
2026-09-25 12:11:00 +03:00
e9f824d8c0 llama : add llama_prec_policy + model-driven W4A4 path (#24364)
* Rebase and update based on #26675

Signed-off-by: ynankani <[email protected]>

* CI failure fix(launh_bounds overload on HIP) and cleanup

Signed-off-by: ynankani <[email protected]>

* Address review comments

Signed-off-by: ynankani <[email protected]>

* Use ggml tensor instead of name in act policy map

Signed-off-by: ynankani <[email protected]>

* Address review comments and cleanup

Signed-off-by: ynankani <[email protected]>

* Address review comments

Signed-off-by: ynankani <[email protected]>

* Rename changes

Signed-off-by: ynankani <[email protected]>

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

Co-authored-by: Georgi Gerganov <[email protected]>

* MXFP4 dispatch changes for higher src prec

Signed-off-by: ynankani <[email protected]>

* Refactor and address review comments

Signed-off-by: ynankani <[email protected]>

* Updates based on review comments

Signed-off-by: ynankani <[email protected]>

* Apply batched suggestions from code review

Co-authored-by: Johannes Gäßler <[email protected]>

* Address review comments

Signed-off-by: ynankani <[email protected]>

* Apply patch from review

Signed-off-by: ynankani <[email protected]>

---------

Signed-off-by: ynankani <[email protected]>
Co-authored-by: Georgi Gerganov <[email protected]>
Co-authored-by: Johannes Gäßler <[email protected]>
2026-09-25 11:36:35 +03:00
uvos d028c697b5 HIP: bump HIP_VERSION requried for fp8 to avoid missing __hip_fp8_e4m3 support in 6.2 (#29231) 2026-09-25 10:36:36 +03:00
Jess SullivanandGeorgi Gerganov 66963a8bc7 rpc: include nb in the get_alloc_size cache key and floor the result at ggml_nbytes (#29283)
* rpc : include nb in the get_alloc_size cache key and floor the result at ggml_nbytes

* cont : remove redundant comment

* cont : add TODO

---------

Co-authored-by: Georgi Gerganov <[email protected]>
2026-09-25 10:35:09 +03:00
Neo Zhang cd74ef6274 [SYCL] support sparse FA (#28796)
* fix conflict

* fix format issue

* rm unused code
2026-09-25 10:17:51 +03:00
R0CKSTAR f9af9be219 musa: fix PH1 (MTT S5000) operator failures and build issues (#29193)
* musa: use 16-byte copies for MUSA like sm_70+

ggml_cuda_get_max_cpy_bytes() derives the copy width from __CUDA_ARCH__. mcc
never defines it, so MUSA fell into the generic branch and returned 8 bytes
instead of the 16 bytes that every sm_70+ target gets. The value sizes the
per-thread copy unit of the FlashAttention K/V staging code (fattn-common,
fattn-vec, fattn-tile, fattn-mma-f16 shared-memory loads) and of mmq-vec-dot,
so every MUSA FlashAttention kernel moved half as many bytes per instruction.

On an MTT S5000 (mp_31, MUSA SDK 5.2.0) with Qwen3.8-27B-UD-Q4_K_M, -ngl 999,
-p 512 -n 64, -fa on: 751.15 -> 794.73 t/s prefill and 15.59 -> 15.69 t/s
decode. -fa off is unchanged (1050.05 -> 1052.86 t/s prefill), FLASH_ATTN_EXT
is unchanged (3984 ok / 0 fail / 1323 unsupported) and perplexity is
unchanged.

* musa: enable the CUB paths on MUSA

GGML_CUDA_USE_CUB and USE_CUB are selected by "CUDART_VERSION >= 11070", which
the MUSA SDK never satisfies: CUDART_VERSION is not defined anywhere under
/usr/local/musa/include, so the condition is always false and every CUB-based
path stayed compiled out on MUSA even though the SDK ships CUB and the kernels
build for mp_31.  Select them from GGML_USE_MUSA as well.  The device-wide
algorithms are usable too: cub::DeviceSegmentedSort compiles and produces
correct results on mp_31.

This lifts the ne[0] <= 1024 limit that ggml_backend_cuda_device_supports_op
applied to ARGSORT and TOP_K on MUSA.  On an MTT S5000 (S5000, mcc 5.2.0):
ARGSORT 48 ok / 52 not supported -> 100 ok / 0 (CUDA parity), TOP_K 0 ok /
354 not supported -> 527 ok / 0.  The other 20 per-op suites are unchanged, the
Qwen3-0.6B f16 (14.4679) and Qwen3.8-27B iq4_nl (5.1724) perplexities are
unchanged, and the 0.6B graph keeps the same nodes and splits (18 CPU + 18
MUSA0, SET_ROWS 1008) as before.

* musa: take the upstream code path where the toolkit supports it

Several guards were written for an older MUSA toolkit. Verified against MUSA SDK
5.2.0 and on an MTT S5000 (mp_31):

- device init: query cudaDevAttrCooperativeLaunch instead of hardcoding false.
  The device reports cooperativeLaunch=1 and musaLaunchCooperativeKernel works
  (verified with a kernel whose result was checked).
- device init: keep prop.warpSize instead of overriding it with 32. The device
  reports 32 anyway, so this only removes the divergence.
- CUDA_SET_SHARED_MEMORY_LIMIT and the FA shared-memory raise: musaFuncSetAttribute
  returns success and sharedMemPerBlockOptin is 192 KiB, so the kernels can use
  more than the default 48 KiB.
- vendors/musa.h: add the cudaDeviceGetAttribute and cudaDevAttrCooperativeLaunch
  mappings the device-init change needs.

Measured on one S5000 with Qwen3.8-27B Q4_K_M (-ngl 999, -r 3): pp512 968.27 ->
957.09 t/s, tg64 10.09 -> 10.23 t/s, FLASH_ATTN_EXT sweep identical (3975/3982
both), perplexity identical (80.2841 +/- 7.26772 both).

* musa: drop compile-time guards that MUSA's runtime gates already cover

mcc never defines __CUDA_ARCH__, so the arch-gated fallbacks in this group
were already taken on MUSA and the GGML_USE_MUSA guards on top of them only
kept the upstream text from being compiled:

  - wkv.cu: the "#pragma unroll" suppression has no effect on the generated
    code that is not already covered by the surrounding guards
  - common.cuh: the MUSA-only __builtin_unreachable() in no_device_code() is
    not needed to silence the compiler
  - ssm-scan.cu: the SSD (Mamba-2 prefill) block and its dispatch are gated at
    runtime by GGML_CUDA_CC_IS_NVIDIA(cc) and turing_mma_available(cc), which
    are both false for PH1 (cc 0x100310), so compiling them changes nothing
  - common.cuh: warp_reduce_max(half2) is guarded the same way as
    warp_reduce_sum(half2) (FP16_AVAILABLE); the MUSA-only guard left the
    function with no return statement. It has no caller today.

MTT S5000 (mp_31, MUSA SDK 5.2.0), MUSA_ARCHITECTURES=31: build rc=0. Against
an unmodified build of the same tree on the same card, FLASH_ATTN_EXT
(3984 ok / 0 fail / 1323 unsupported), SSM_SCAN (15/0), RWKV_WKV6 (6/0),
GATED_DELTA_NET (38/0) and MUL_MAT (1299/0/385 unsupported) are identical, and
perplexity with -fa on is bit-identical (5.1639 +/- 0.36673, 4 chunks).

* musa: do not use MMQ on PH1

test-backend-ops on an MTT S5000 (mp_31, MUSA SDK 5.2.0) fails 260 cases and every
one of them goes through the MMQ path:

  - MUL_MAT with a batched src1 (any bs/nr != [1,1]): 109 cases across all
    quantized types, e.g. 12 of 13 cases at n=16, while the plain [1,1] layout
    passes
  - every quantized MUL_MAT_ID: 147 cases, while the f16/f32 variants of the same
    shapes pass
  - MUL_MAT with more than ~512 tokens: 4 cases (n=509..4096); the small-n cases pass

The cuBLAS/dequant path is correct for all of them and the MMVQ path used for
small batches is unaffected, so quantized matmuls now take that path on PH1
instead of returning wrong values. 27B perplexity with default flags goes from
nan to finite, and the full suite reports 0 failures out of 22237 cases.

The MMQ defect itself (fastdiv, __umulhi, uint3 kernel parameters and
__CUDA_ARCH__-based MMA availability were all checked and are correct on this
part) is not addressed here.

* musa: keep the block barrier of the fused TOPK_MOE kernel reachable

topk_moe_cuda returns early for the rows past the end of the graph, but one block
covers TOPK_MOE_ROWS_PER_BLOCK (8) rows, so the last block is only partially filled
whenever n_rows is not a multiple of 8.  On MUSA a warp that has already returned
blocks the block wide __syncthreads() below, which makes the kernel hang and the
launch time out.  CUDA tolerates the exited warps, which is why the CUDA numbers
never showed it.

For MUSA, clamp the row index of those warps to the last row so that every warp of
the block reaches the barrier; they recompute the last row and write the same
values.  The CUDA code path is unchanged.

On an MTT S5000 (mp_31) the fused TOPK_MOE cases change from a launch timeout with
no completed case to 418 ok / 0 not supported / 0 failed, i.e. the CUDA result, and
the other 101 per op suites are unchanged (0 failed, no count changes).

* musa: enable GATED_DELTA_NET

The op was turned off for every MUSA target because mcc could not build the kernel
at the time. The current toolkit builds it: with mp_31 and MUSA SDK 5.2.0 the file
compiles with zero errors and all 36 test-backend-ops GATED_DELTA_NET cases pass
against the CPU reference. 27B perplexity is unchanged.

While the op is refused, the scheduler has no choice but to run it on the CPU: 48
GATED_DELTA_NET nodes per forward pass. On an MTT S5000 (Qwen3.8-27B Q4_K_M, -ngl
999, one container, -r 3):

    pp512 (FA off)   964.51 -> 2119.26 t/s
    tg64  (FA off)    10.15 ->   15.50 t/s

* musa: name the stream capture query API for the graph aware kernels

argsort.cu and mean.cu call cudaStreamCaptureStatus, cudaStreamIsCapturing and
cudaStreamCaptureStatusNone inside their USE_CUDA_GRAPH blocks, but the MUSA
compatibility headers do not alias those names, so building with the experimental
GGML_MUSA_GRAPHS option fails with 7 errors in those two files.  Map the three
names to their musa* counterparts, under the same guard that enables the graph
code, so the default build is untouched.

The option stays off by default: on an MTT S5000 the captured path measured
slower (pp512 693 vs 772 t/s, tg128 15.20 vs 15.39 t/s over two sessions) and the
borderline MUL_MAT cases are not reproducible between runs.

* musa: build the CI and docs for PH1 (MTT S5000)

The MUSA CI job and the documented default still targeted the first generation
(MTT S80, MUSA_ARCHITECTURES=21) while the current MUSA SDK targets PH1
(MTT S5000, 31).  Move the job, ci/run.sh's default and the build docs to 31,
and run the job in the PH1 MUSA SDK devel image:

    registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu22.04-amd64

That image needs two things the previous one did not: python3-venv for the
ccache-buckets step, which builds a virtual environment for the Hugging Face
CLI, and no time prefix on the build command, because container jobs run their
steps with sh and the image ships no time binary.
2026-09-25 10:16:29 +03:00
InflexCZE 1ab7e5ad2d CUDA: fuse RMS_NORM + SCALE into one kernel (#29393)
- #28068 builds the GDN q/k l2norm as ggml_scale(ggml_rms_norm(x, eps/n), 1/sqrt(n)). This adds 2 SCALE nodes per GDN layer, 96 extra kernel launches per ubatch on Qwen3.8-27B (48 GDN layers).
- The extra kernels take no measurable GPU time, but each launch has a host/driver cost. It is small with plain batch processing and about 10x larger with draft-mtp speculative decoding.
- rms_norm_f32 gets a do_scale flag, the same pattern as do_multiply/do_add, so the fused path shares the kernel, the reduction and the launcher. It computes scale * (rsqrt(mean + eps) * x), which matches the unfused rms_norm + scale bit for bit, so #28068 numerics are kept.
- Fusion only fires when SCALE has no bias and the rms_norm output has a single consumer (ggml_can_fuse).
- Metal (#28948) and SYCL (#28931) already fuse the same pattern.

Measured on 2x GTX 1080 Ti (sm_61, PCIe 3.0 x16 + x4), i7-13700KF, Windows 11, driver 582.66, CUDA 12.9.
Qwen3.8-27B-UD-Q4_K_XL, -ngl 99 -ts 53,47 -ot token_embd=CPU, master fee39dd92.

llama-bench -ub 128,512 -p 512,2048 -n 128 -r 5, tok/s:

  build           pp512@128  pp2048@128  pp2048@512  tg128
  master           367.7      419.1       385.4      12.90
  master + fix     372.6      420.6       388.6      12.98
                   +1.3%      +0.4%       +0.8%      +0.6%

llama-server cold prefill, -c 56000 -ub 128 -b 2048, draft-mtp n-max 3 p-min 0.5, mean of 2 rounds x 3 reps:

  build           pp 8000         pp 20000
  master          356.5           322.0
  master + fix    371.4 (+4.2%)   337.4 (+4.8%)

- Launches per ubatch go from 1032.9 + 841.7 back to 978.9 + 799.7 (CUDA0 + CUDA1), the b10828 count. The GPU op sum is unchanged.
- test-backend-ops RMS_NORM_SCALE, NORM_SCALE, RMS_NORM_MUL_ADD, RMS_NORM_MUL_ROPE, RMS_NORM, RMS_NORM_BACK, NORM, L2_NORM and SCALE all pass on both GPUs.
- Perplexity is identical to the unfused build: 3.2030 +/- 0.0559 at -c 2048, 16 chunks.
- Draft acceptance counts per request match the unfused build.

Assisted-by: Claude Opus 5.5
2026-09-25 08:22:43 +03:00
Aman Gupta f805c57a2d llama : fix tensor split for fused qkv with uneven K/V head sizes (#29294)
* llama : fix tensor split for fused qkv with uneven K/V head sizes

Assisted-by: Qwen3.8-27B

* fix v granularity

* convert: fix mtp conversion

* convert: add support for mtp flags

* fix loader
2026-09-25 11:15:33 +08:00
Jhen-Jie Hong 4de0926596 hexagon: add q5_k quant type support (#29123) 2026-09-24 19:03:59 -07:00
kurquhar ed319febb1 hexagon: use DMA for contiguous dim1 CONCAT (#29404)
* hexagon: use DMA for contiguous dim1 CONCAT

Assisted-by: OpenCode

* hexagon: update CONCAT DMA for DMA64

Assisted-by: OpenCode
2026-09-24 19:03:41 -07:00
Georgi Gerganov 84e76d8a23 metal : fix graph capture and handle empty graphs (#29390)
- return early when the graph has no nodes
- drop the redundant reset of capture_compute: the decrement at the top
  of the function already transitions the counter from 0 to -1, so a
  capture happens exactly once
- hint at METAL_CAPTURE_ENABLED=1 in the capture error message
- pass capture_compute == 0 (not the raw counter) as use_capture to
  ggml_metal_op_init, so GPU debug-group markers are only emitted on the
  captured compute

Assisted-by: pi:llama.cpp/Qwen3.8-27B
2026-09-24 22:44:53 +03:00
Georgi Gerganov cdc06426e7 metal : optimize sparse FA + clean-up (#29377)
* metal : cache sparse FA indices in shared memory

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

* metal : simplify shared memory size calculation

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

* pi : update general

* metal : unroll sparse index load

Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
2026-09-24 22:44:36 +03:00
Georgi Gerganov bced4595b8 sync : ggml (#29396)
* ggml : bump version to 0.25.2 (ggml/1642)

* ggml : fix ubsan error in `ggml_graph_nbytes` (ggml/1644)

* ggml : bump version to 0.25.3 (ggml/1645)

* sync : ggml
2026-09-24 22:44:05 +03:00
Jhen-Jie Hong a02c7f58c1 hexagon: handle multi-sequence in concat_2d (#29344) 2026-09-24 12:16:55 -07:00
Daniel Kuts 5cf3a35287 llama-grammar: fix numeric truncation for token_id parsing (#29382) 2026-09-24 21:43:26 +03:00
07fc586e38 hexagon: dynamic quantizer improvements (#29395)
* hexagon: fix accuracy issue in Q8_0 N=1 MUL_MAT

* hex-quant: fix register spills

* hex-mm: use dma for all dyn.quant paths

Co-authored-by: Aparna M P <[email protected]>

* hex-mm: remove obsolete run_quant_task

* hex-mm: update tracing to properly wrap the events

* hex-mm: use act for activation data in all paths

* hex-mm: use act_ instead of src1_ to avoid confusion in fused kernels

* hex-mm: remove/reroute the rest of the non-DMA act (aka src1) logic

* hex-dma64: yet another pass at cleaning up the dma_addr_t casts

* Update ggml/src/ggml-hexagon/htp/matmul-ops.h

Co-authored-by: Sigbjørn Skjæret <[email protected]>

* Update ggml/src/ggml-hexagon/htp/matmul-ops.c

Co-authored-by: Sigbjørn Skjæret <[email protected]>

* Update ggml/src/ggml-hexagon/htp/matmul-ops.c

Co-authored-by: Sigbjørn Skjæret <[email protected]>

* Update ggml/src/ggml-hexagon/htp/matmul-ops.c

Co-authored-by: Sigbjørn Skjæret <[email protected]>

---------

Co-authored-by: Max Krasnyansky <[email protected]>
Co-authored-by: Aparna M P <[email protected]>
Co-authored-by: Sigbjørn Skjæret <[email protected]>
2026-09-24 11:39:28 -07:00
222 changed files with 24803 additions and 17041 deletions
+3 -3
View File
@@ -1,4 +1,4 @@
ARG ONEAPI_VERSION=2025.3.3-0-devel-ubuntu24.04
ARG ONEAPI_VERSION=2026.1.1-devel-ubuntu24.04
ARG BUILD_DATE=N/A
ARG APP_VERSION=N/A
ARG APP_REVISION=N/A
@@ -19,7 +19,7 @@ RUN npm ci
COPY tools/ui/ ./
RUN LLAMA_BUILD_NUMBER="$APP_VERSION" npm run build
FROM docker.io/intel/deep-learning-essentials:$ONEAPI_VERSION AS build
FROM docker.io/intel/oneapi-toolkit:$ONEAPI_VERSION AS build
ARG GGML_SYCL_F16=ON
ARG LEVEL_ZERO_VERSION=1.28.2
@@ -59,7 +59,7 @@ RUN mkdir -p /app/full \
&& cp requirements.txt /app/full \
&& cp .devops/tools.sh /app/full/tools.sh
FROM docker.io/intel/deep-learning-essentials:$ONEAPI_VERSION AS base
FROM docker.io/intel/oneapi-toolkit:$ONEAPI_VERSION AS base
ARG BUILD_DATE=N/A
ARG APP_VERSION=N/A
+2 -3
View File
@@ -1,10 +1,9 @@
ARG UBUNTU_VERSION=22.04
# This needs to generally match the container host's environment.
ARG MUSA_VERSION=rc4.3.0
# Target the MUSA build image
ARG BASE_MUSA_DEV_CONTAINER=docker.io/mthreads/musa:${MUSA_VERSION}-devel-ubuntu${UBUNTU_VERSION}-amd64
ARG BASE_MUSA_DEV_CONTAINER=registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu${UBUNTU_VERSION}-amd64
ARG BASE_MUSA_RUN_CONTAINER=docker.io/mthreads/musa:${MUSA_VERSION}-runtime-ubuntu${UBUNTU_VERSION}-amd64
ARG BASE_MUSA_RUN_CONTAINER=${BASE_MUSA_DEV_CONTAINER}
ARG BUILD_DATE=N/A
ARG APP_VERSION=N/A
+3 -1
View File
@@ -33,6 +33,7 @@ concurrency:
env:
GGML_NLOOP: 3
GGML_N_THREADS: 1
GGML_SCHED_DEBUG_REALLOC: 1
LLAMA_ARG_LOG_COLORS: 1
LLAMA_ARG_LOG_PREFIX: 1
LLAMA_ARG_LOG_TIMESTAMPS: 1
@@ -98,7 +99,8 @@ jobs:
id: cmake_test
run: |
cd build
ctest -L main -E "test-llama-archs" --verbose --timeout 900
# ref: https://github.com/ggml-org/llama.cpp/pull/19802#issuecomment-4013704023
ctest -L main -E "test-llama-archs|test-save-load-state" --verbose --timeout 900
macos-latest-x64:
runs-on: macos-15-intel
+1
View File
@@ -37,6 +37,7 @@ concurrency:
env:
GGML_NLOOP: 3
GGML_N_THREADS: 1
GGML_SCHED_DEBUG_REALLOC: 1
LLAMA_ARG_LOG_COLORS: 1
LLAMA_ARG_LOG_PREFIX: 1
LLAMA_ARG_LOG_TIMESTAMPS: 1
+4 -4
View File
@@ -145,7 +145,7 @@ jobs:
musa:
runs-on: ubuntu-22.04
container: mthreads/musa:rc4.3.0-devel-ubuntu22.04-amd64
container: registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu22.04-amd64
steps:
- name: Clone
@@ -156,7 +156,7 @@ jobs:
id: depends
run: |
apt-get update
apt-get install -y build-essential git cmake libssl-dev jq
apt-get install -y build-essential git cmake libssl-dev jq python3-venv
- name: ccache
uses: ggml-org/[email protected]
@@ -178,8 +178,8 @@ jobs:
run: |
cmake -B build -S . \
-DGGML_MUSA=ON \
-DMUSA_ARCHITECTURES=21
time cmake --build build --config Release -j $(nproc)
-DMUSA_ARCHITECTURES=31
cmake --build build --config Release -j $(nproc)
- name: ccache-buckets-save
if: ${{ github.event_name == 'push' && github.ref == 'refs/heads/master' }}
+5 -5
View File
@@ -48,7 +48,7 @@ jobs:
env:
ONEAPI_ROOT: /opt/intel/oneapi/
ONEAPI_INSTALLER_VERSION: "2025.3.3"
ONEAPI_INSTALLER_VERSION: "2026.1"
LEVEL_ZERO_VERSION: "1.33.1"
LEVEL_ZERO_UBUNTU_VERSION: "u24.04"
@@ -63,8 +63,8 @@ jobs:
shell: bash
run: |
cd /tmp
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/56f7923a-adb8-43f3-8b02-2b60fcac8cab/intel-deep-learning-essentials-2025.3.3.16_offline.sh -O intel-deep-learning-essentials_offline.sh
sudo bash intel-deep-learning-essentials_offline.sh -s -a --silent --eula accept
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/5996e26b-f48a-42b1-8db0-b002ad0bd8d7/intel-oneapi-toolkit-2026.1.1.33_offline.sh -O intel-oneapi-toolkit_offline.sh
sudo bash intel-oneapi-toolkit_offline.sh -s -a --silent --eula accept
- name: Install Level Zero SDK
shell: bash
@@ -129,11 +129,11 @@ jobs:
shell: bash
env:
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/b60765d1-2b85-4e85-86b6-cb0e9563a699/intel-deep-learning-essentials-2025.3.3.18_offline.exe
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/0cb67a0d-67f6-410b-868b-f4a0a17ff0cf/intel-oneapi-toolkit-2026.1.1.32_offline.exe
WINDOWS_DPCPP_MKL: intel.oneapi.win.cpp-dpcpp-common:intel.oneapi.win.mkl.devel:intel.oneapi.win.dnnl:intel.oneapi.win.tbb.devel
LEVEL_ZERO_SDK_URL: https://github.com/oneapi-src/level-zero/releases/download/v1.33.1/level-zero-win-sdk-1.33.1.zip
ONEAPI_ROOT: "C:/Program Files (x86)/Intel/oneAPI"
ONEAPI_INSTALLER_VERSION: "2025.3.3"
ONEAPI_INSTALLER_VERSION: "2026.1"
steps:
- name: Clone
id: checkout
+1
View File
@@ -31,6 +31,7 @@ concurrency:
env:
GGML_NLOOP: 3
GGML_N_THREADS: 1
GGML_SCHED_DEBUG_REALLOC: 1
LLAMA_ARG_LOG_COLORS: 1
LLAMA_ARG_LOG_PREFIX: 1
LLAMA_ARG_LOG_TIMESTAMPS: 1
+1 -1
View File
@@ -31,7 +31,7 @@ jobs:
uses: actions/setup-python@v6
with:
python-version: "3.11"
pip-install: -r requirements/requirements-all.txt ty==0.0.78
pip-install: -r requirements/requirements-all.txt ty==0.0.84
# - name: Type-check with Pyright
# uses: jakebailey/pyright-action@v2
# with:
+16 -16
View File
@@ -1344,11 +1344,11 @@ jobs:
shell: bash
env:
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/b60765d1-2b85-4e85-86b6-cb0e9563a699/intel-deep-learning-essentials-2025.3.3.18_offline.exe
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/0cb67a0d-67f6-410b-868b-f4a0a17ff0cf/intel-oneapi-toolkit-2026.1.1.32_offline.exe
WINDOWS_DPCPP_MKL: intel.oneapi.win.cpp-dpcpp-common:intel.oneapi.win.mkl.devel:intel.oneapi.win.dnnl:intel.oneapi.win.tbb.devel
LEVEL_ZERO_SDK_URL: https://github.com/oneapi-src/level-zero/releases/download/v1.28.2/level-zero-win-sdk-1.28.2.zip
LEVEL_ZERO_SDK_URL: https://github.com/oneapi-src/level-zero/releases/download/v1.33.1/level-zero-win-sdk-1.33.1.zip
ONEAPI_ROOT: "C:/Program Files (x86)/Intel/oneAPI"
ONEAPI_INSTALLER_VERSION: "2025.3.3"
ONEAPI_INSTALLER_VERSION: "2026.1"
steps:
- name: Clone
@@ -1391,9 +1391,11 @@ jobs:
run: |
echo "cp oneAPI running time dll files in ${{ env.ONEAPI_ROOT }} to ./build/bin"
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_sycl_blas.5.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_core.2.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_tbb_thread.2.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_sycl_blas.6.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_core.3.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_def.3.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_avx2.3.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_tbb_thread.3.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/ur_adapter_level_zero.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/ur_adapter_level_zero_v2.dll" ./build/bin
@@ -1408,13 +1410,11 @@ jobs:
echo "Level Zero loader DLL not found in oneAPI or SDK; relying on system driver/runtime"
fi
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/sycl8.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/sycl9.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/svml_dispmd.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libmmd.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libiomp5md.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/sycl-ls.exe" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libsycl-fallback-bfloat16.spv" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libsycl-native-bfloat16.spv" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/dnnl/latest/bin/dnnl.dll" ./build/bin
cp "${{ env.ONEAPI_ROOT }}/tbb/latest/bin/tbb12.dll" ./build/bin
@@ -1454,8 +1454,8 @@ jobs:
env:
ONEAPI_ROOT: /opt/intel/oneapi/
ONEAPI_INSTALLER_VERSION: "2025.3.3"
LEVEL_ZERO_VERSION: "1.28.2"
ONEAPI_INSTALLER_VERSION: "2026.1"
LEVEL_ZERO_VERSION: "1.33.1"
LEVEL_ZERO_UBUNTU_VERSION: "u24.04"
steps:
@@ -1469,16 +1469,16 @@ jobs:
shell: bash
run: |
cd /tmp
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/56f7923a-adb8-43f3-8b02-2b60fcac8cab/intel-deep-learning-essentials-2025.3.3.16_offline.sh -O intel-deep-learning-essentials_offline.sh
sudo bash intel-deep-learning-essentials_offline.sh -s -a --silent --eula accept
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/5996e26b-f48a-42b1-8db0-b002ad0bd8d7/intel-oneapi-toolkit-2026.1.1.33_offline.sh -O intel-oneapi-toolkit_offline.sh
sudo bash intel-oneapi-toolkit_offline.sh -s -a --silent --eula accept
- name: Install Level Zero SDK
shell: bash
run: |
cd /tmp
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/level-zero_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O level-zero.deb
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/level-zero-devel_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O level-zero-devel.deb
sudo apt-get install -y ./level-zero.deb ./level-zero-devel.deb
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/libze1_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O libze1.deb
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/libze-dev_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O libze-dev.deb
sudo apt-get install -y ./libze1.deb ./libze-dev.deb
- name: Download UI build
uses: actions/download-artifact@v7
+1
View File
@@ -7,6 +7,7 @@ General:
- Don't try to build or run the code unless you are explicitly asked to do so
- Use the `gh` CLI tool when querying PRs, issues, or other GitHub resources
- When [MODEL] is needed, first try to get it from the `PI_MODEL_NAME` env var before asking the user
- Never read the `AGENTS.md` file
Coding:
- When in doubt, always refer to the CONTRIBUTING.md file of the project
+1 -1
View File
@@ -57,7 +57,7 @@
/ggml/src/ggml-cann/ @ggml-org/ggml-cann
/ggml/src/ggml-common.h @ggerganov
/ggml/src/ggml-cpu/ @ggerganov
/ggml/src/ggml-cpu/iqp.* @bartowski1182
/ggml/src/ggml-cpu/tiled/ @jbooth @bartowski1182
/ggml/src/ggml-cpu/spacemit/ @alex-spacemit
/ggml/src/ggml-cuda/ @ggml-org/ggml-cuda
/ggml/src/ggml-cuda/vendors/hip.h @IMbackK
+1 -1
View File
@@ -21,7 +21,7 @@ docker run --privileged -it \
-v $HOME/llama.cpp/ci-cache:/ci-cache \
-v $HOME/llama.cpp/ci-results:/ci-results \
-v $PWD:/ws -w /ws \
mthreads/musa:rc4.3.0-devel-ubuntu22.04-amd64
registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu22.04-amd64
```
Inside the container, execute the following commands:
+2 -2
View File
@@ -158,8 +158,8 @@ if [ ! -z ${GG_BUILD_WEBGPU} ]; then
fi
if [ ! -z ${GG_BUILD_MUSA} ]; then
# Use qy1 by default (MTT S80)
MUSA_ARCH=${MUSA_ARCH:-21}
# Use ph1 by default (MTT S5000)
MUSA_ARCH=${MUSA_ARCH:-31}
CMAKE_EXTRA="${CMAKE_EXTRA} -DGGML_MUSA=ON -DMUSA_ARCHITECTURES=${MUSA_ARCH}"
fi
+11 -10
View File
@@ -351,7 +351,7 @@ static bool parse_bool_value(const std::string & value) {
static std::string get_default_local_path(const std::string & url) {
auto f = string_split<std::string>(url, '#').front();
f = string_split<std::string>(f, '?').front();
return fs_get_cache_file(string_split<std::string>(f, '/').back());
return fs_path_to_utf8(fs_get_cache_file(string_split<std::string>(f, '/').back()));
}
static bool spec_types_is_default(const common_params & params) {
@@ -2673,16 +2673,17 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
params.video_ffmpeg_bin_dir = value;
}
).set_examples(mmproj_examples).set_env("LLAMA_ARG_VIDEO_FFMPEG_DIR"));
if (params.is_gen_docs || llama_supports_rpc()) {
add_opt(common_arg(
{"--rpc"}, "SERVERS",
"comma-separated list of RPC servers (host:port)",
[](common_params & params, const std::string & value) {
add_rpc_devices(value);
GGML_UNUSED(params);
add_opt(common_arg(
{"--rpc"}, "SERVERS",
"comma-separated list of RPC servers (host:port)",
[](common_params & params, const std::string & value) {
if (!llama_supports_rpc()) {
throw std::invalid_argument("RPC not supported in this build");
}
).set_env("LLAMA_ARG_RPC"));
}
add_rpc_devices(value);
GGML_UNUSED(params);
}
).set_env("LLAMA_ARG_RPC"));
add_opt(common_arg(
{"-lm", "--load-mode"}, "MODE",
"model loading mode (default: auto)\n"
+206 -179
View File
@@ -46,11 +46,10 @@
#include <io.h>
#else
#include <sys/ioctl.h>
#include <sys/stat.h>
#include <unistd.h>
#endif
#if defined(__linux__)
#if !defined(_WIN32) && !defined(__APPLE__)
#include <sys/types.h>
#include <pwd.h>
#endif
@@ -614,34 +613,6 @@ std::string string_from(const struct llama_context * ctx, const std::vector<llam
return buf.str();
}
std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch) {
std::stringstream buf;
buf << "[ ";
bool first = true;
for (int i = 0; i < batch.n_tokens; ++i) {
if (!first) {
buf << ", ";
} else {
first = false;
}
auto detokenized = common_token_to_piece(ctx, batch.token[i]);
buf << "\n" << std::to_string(i)
<< ", token '" << detokenized << "'"
<< ", pos " << std::to_string(batch.pos[i])
<< ", n_seq_id " << std::to_string(batch.n_seq_id[i])
<< ", seq_id " << std::to_string(batch.seq_id[i][0])
<< ", logits " << std::to_string(batch.logits[i]);
}
buf << " ]";
return buf.str();
}
void string_process_escapes(std::string & input) {
std::size_t input_len = input.length();
std::size_t output_idx = 0;
@@ -900,7 +871,7 @@ bool fs_validate_filename(const std::string & filename, bool allow_subdirs) {
#ifdef _WIN32
static std::wstring utf8_to_wstring(const std::string & str) {
std::wstring utf8_to_wstring(const std::string & str) {
if (str.empty()) {
return std::wstring();
}
@@ -916,82 +887,29 @@ static std::wstring utf8_to_wstring(const std::string & str) {
return wstr;
}
std::string wstring_to_utf8(const std::wstring & str) {
if (str.empty()) {
return std::string();
}
int size = WideCharToMultiByte(CP_UTF8, 0, str.c_str(), (int)str.size(), NULL, 0, NULL, NULL);
if (size <= 0) {
return std::string();
}
std::string utf8(size, 0);
WideCharToMultiByte(CP_UTF8, 0, str.c_str(), (int)str.size(), &utf8[0], size, NULL, NULL);
return utf8;
}
#endif
// returns true if successful, false otherwise
bool fs_create_directory_with_parents(const std::string & path) {
#ifdef _WIN32
std::wstring wpath = utf8_to_wstring(path);
// if the path already exists, check whether it's a directory
const DWORD attributes = GetFileAttributesW(wpath.c_str());
if ((attributes != INVALID_FILE_ATTRIBUTES) && (attributes & FILE_ATTRIBUTE_DIRECTORY)) {
return true;
}
size_t pos_slash = 0;
// process path from front to back, procedurally creating directories
while ((pos_slash = path.find('\\', pos_slash)) != std::string::npos) {
const std::wstring subpath = wpath.substr(0, pos_slash);
pos_slash += 1;
// skip the drive letter, in some systems it can return an access denied error
if (subpath.length() == 2 && subpath[1] == ':') {
continue;
}
const bool success = CreateDirectoryW(subpath.c_str(), NULL);
if (!success) {
const DWORD error = GetLastError();
// if the path already exists, ensure that it's a directory
if (error == ERROR_ALREADY_EXISTS) {
const DWORD attributes = GetFileAttributesW(subpath.c_str());
if (attributes == INVALID_FILE_ATTRIBUTES || !(attributes & FILE_ATTRIBUTE_DIRECTORY)) {
return false;
}
} else {
return false;
}
}
}
return true;
#else
// if the path already exists, check whether it's a directory
struct stat info;
if (stat(path.c_str(), &info) == 0) {
return S_ISDIR(info.st_mode);
}
size_t pos_slash = 1; // skip leading slashes for directory creation
// process path from front to back, procedurally creating directories
while ((pos_slash = path.find('/', pos_slash)) != std::string::npos) {
const std::string subpath = path.substr(0, pos_slash);
struct stat info;
// if the path already exists, ensure that it's a directory
if (stat(subpath.c_str(), &info) == 0) {
if (!S_ISDIR(info.st_mode)) {
return false;
}
} else {
// create parent directories
const int ret = mkdir(subpath.c_str(), 0755);
if (ret != 0) {
return false;
}
}
pos_slash += 1;
}
return true;
#endif // _WIN32
// returns the path as a UTF-8 string, preserving its separators
std::string fs_path_to_utf8(const std::filesystem::path & path) {
const auto value = path.u8string();
return std::string(value.begin(), value.end());
}
bool fs_is_directory(const std::string & path) {
@@ -1016,58 +934,52 @@ void common_set_env(const std::string & name, const std::string & value) {
#endif
}
std::string fs_get_cache_directory() {
std::string cache_directory = "";
auto ensure_trailing_slash = [](std::string p) {
// Make sure to add trailing slash
if (p.empty() || p.back() != DIRECTORY_SEPARATOR) {
p += DIRECTORY_SEPARATOR;
}
return p;
};
cache_directory = common_get_env("LLAMA_CACHE");
std::filesystem::path common_get_path_from_env(const std::string & name) {
#if defined(_WIN32)
const std::wstring wname = utf8_to_wstring(name);
const wchar_t * wvalue = _wgetenv(wname.c_str());
return wvalue ? std::filesystem::path(wvalue) : std::filesystem::path();
#else
const char * value = std::getenv(name.c_str());
return value ? std::filesystem::path(value) : std::filesystem::path();
#endif
}
std::filesystem::path fs_get_cache_directory() {
std::filesystem::path cache_directory = common_get_path_from_env("LLAMA_CACHE");
if (!cache_directory.empty()) {
return cache_directory;
}
#if defined(_WIN32)
cache_directory = common_get_path_from_env("LOCALAPPDATA");
if (cache_directory.empty()) {
#if defined(__linux__) || defined(__FreeBSD__) || defined(_AIX) || \
defined(__OpenBSD__) || defined(__NetBSD__)
const std::string xdg_cache_home = common_get_env("XDG_CACHE_HOME");
const std::string home = common_get_env("HOME");
if (!xdg_cache_home.empty()) {
cache_directory = xdg_cache_home;
} else if (!home.empty()) {
cache_directory = home + "/.cache/";
throw std::runtime_error("Failed to find %LOCALAPPDATA% directory");
}
#elif defined(__APPLE__)
cache_directory = common_get_path_from_env("HOME");
if (cache_directory.empty()) {
throw std::runtime_error("Failed to find $HOME directory");
}
cache_directory /= "Library/Caches";
#else
cache_directory = common_get_path_from_env("XDG_CACHE_HOME");
if (cache_directory.empty()) {
cache_directory = common_get_path_from_env("HOME");
if (!cache_directory.empty()) {
cache_directory /= ".cache";
} else {
#if defined(__linux__)
/* no $HOME is defined, fallback to getpwuid */
struct passwd *pw = getpwuid(getuid());
if ((!pw) || (!pw->pw_dir)) {
const struct passwd * pw = getpwuid(getuid());
if (!pw || !pw->pw_dir || !*pw->pw_dir) {
throw std::runtime_error("Failed to find $HOME directory");
}
cache_directory = std::string(pw->pw_dir) + std::string("/.cache/");
#else /* defined(__linux__) */
throw std::runtime_error("Failed to find $HOME directory");
#endif /* defined(__linux__) */
cache_directory = pw->pw_dir;
cache_directory /= ".cache";
}
#elif defined(__APPLE__)
cache_directory = common_get_env("HOME");
if (cache_directory.empty()) {
throw std::runtime_error("Failed to find $HOME directory");
}
cache_directory += "/Library/Caches/";
#elif defined(_WIN32)
cache_directory = common_get_env("LOCALAPPDATA");
if (cache_directory.empty()) {
throw std::runtime_error("Failed to find %LOCALAPPDATA% directory");
}
#elif defined(__EMSCRIPTEN__)
GGML_ABORT("not implemented on this platform");
#else
# error Unknown architecture
#endif
cache_directory = ensure_trailing_slash(cache_directory);
cache_directory += "llama.cpp";
}
return ensure_trailing_slash(cache_directory);
#endif
return cache_directory / "llama.cpp";
}
std::string fs_get_config_directory() {
@@ -1115,14 +1027,15 @@ std::string fs_get_config_directory() {
return ensure_trailing_slash(config_directory);
}
std::string fs_get_cache_file(const std::string & filename) {
std::filesystem::path fs_get_cache_file(const std::string & filename) {
GGML_ASSERT(filename.find(DIRECTORY_SEPARATOR) == std::string::npos);
std::string cache_directory = fs_get_cache_directory();
const bool success = fs_create_directory_with_parents(cache_directory);
if (!success) {
throw std::runtime_error("failed to create cache directory: " + cache_directory);
const std::filesystem::path cache_directory = fs_get_cache_directory();
std::error_code ec;
std::filesystem::create_directories(cache_directory, ec);
if (ec) {
throw std::runtime_error("failed to create cache directory: " + fs_path_to_utf8(cache_directory));
}
return cache_directory + filename;
return cache_directory / std::filesystem::u8path(filename);
}
std::vector<common_file_info> fs_list(const std::string & path, bool include_directories) {
@@ -1527,7 +1440,8 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
}
if (llama_model_has_encoder(model)) {
llama_encode(lctx, llama_batch_get_one(tmp.data(), tmp.size()));
common_batch batch = common_batch_get_one(lctx, tmp);
llama_process(lctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get());
llama_token decoder_start_token_id = llama_model_decoder_start_token(model);
if (decoder_start_token_id == LLAMA_TOKEN_NULL) {
decoder_start_token_id = bos;
@@ -1536,7 +1450,9 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
tmp.push_back(decoder_start_token_id);
}
if (llama_model_has_decoder(model)) {
llama_decode(lctx, llama_batch_get_one(tmp.data(), std::min(tmp.size(), (size_t) params.n_batch)));
tmp.resize(std::min(tmp.size(), (size_t) params.n_batch));
common_batch batch = common_batch_get_one(lctx, tmp);
llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
llama_memory_clear(llama_get_memory(lctx), true);
llama_synchronize(lctx);
@@ -1600,9 +1516,13 @@ common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx) {
tmp.push_back(0);
tmp.push_back(0);
int ret = llama_decode(ctx, llama_batch_get_one(tmp.data(), tmp.size()));
int ret;
{
common_batch batch = common_batch_get_one(ctx, tmp);
ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
if (ret != 0) {
COM_ERR("llama_decode() failed: %d\n", ret);
COM_ERR("llama_process() failed: %d\n", ret);
res = COMMON_CONTEXT_SEQ_RM_TYPE_NO;
goto done;
}
@@ -2189,29 +2109,138 @@ float lr_opt::get_lr(float epoch) const {
}
bool common_replay_last_token(struct llama_context * ctx, llama_token last_token, int32_t pos) {
llama_batch batch = llama_batch_get_one(&last_token, 1);
batch.pos = &pos;
if (llama_decode(ctx, batch)) {
common_batch batch(ctx);
batch.add(last_token, pos, 0, true);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s: failed to replay last token\n", __func__);
return false;
}
return true;
}
llama_batch_ext_ptr common_batch_ext_get_one(llama_context * ctx, const llama_tokens & tokens) {
llama_batch_ext_ptr batch(llama_batch_ext_init(ctx));
common_batch::common_batch(llama_context * ctx) : batch(llama_batch_ext_init(ctx)) {
const auto rope_type = llama_model_rope_type(llama_get_model(ctx));
n_pos = rope_type == LLAMA_ROPE_TYPE_MROPE || rope_type == LLAMA_ROPE_TYPE_IMROPE ? GGML_MROPE_SECTIONS : 1;
}
auto mem = llama_get_memory(ctx);
llama_pos pos = mem ? llama_memory_seq_pos_max(mem, 0) + 1 : 0;
void common_batch::clear() {
tokens.clear();
llama_batch_ext_clear(batch.get());
}
for (size_t i = 0; i < tokens.size(); ++i) {
const int32_t idx = llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
llama_batch_ext_set_pos(batch.get(), idx, &pos);
pos++;
int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
const int32_t idx = llama_batch_ext_add_token(batch.get(), seq_id, id);
if (idx < 0) {
GGML_ABORT("%s: failed to add token %d to the batch (error %d, n_tokens = %d)\n", __func__, id, idx, size());
}
llama_batch_ext_set_pos(batch.get(), idx, &pos);
if (output) {
llama_batch_ext_set_output_logits(batch.get(), idx, true);
}
tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 } });
return idx;
}
bool common_batch::set_output(int32_t idx, bool value) {
if (idx < 0 || idx >= (int32_t) tokens.size()) {
return false;
}
tokens[idx].output = value;
return llama_batch_ext_set_output_logits(batch.get(), idx, value);
}
bool common_batch::set_embd(int32_t idx, llama_embd embd) {
if (idx < 0 || idx >= (int32_t) tokens.size()) {
return false;
}
if (!llama_batch_ext_set_embd_token(batch.get(), idx, embd)) {
return false;
}
tokens[idx].embd = embd;
return true;
}
int32_t common_batch::add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output) {
const int32_t idx = llama_batch_ext_add_embd(batch.get(), seq_id, embd);
if (idx < 0) {
GGML_ABORT("%s: failed to add embedding to the batch (error %d, n_tokens = %d)\n", __func__, idx, size());
}
llama_batch_ext_set_pos(batch.get(), idx, pos);
if (output) {
llama_batch_ext_set_output_logits(batch.get(), idx, true);
}
token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd };
for (int32_t j = 0; j < n_pos; ++j) {
t.pos[j] = pos[j];
}
tokens.push_back(t);
return idx;
}
common_batch common_batch_from_llama_batch(llama_context * ctx, const llama_batch & batch) {
common_batch res(ctx);
const bool has_token = batch.token != nullptr;
const bool has_embd = batch.embd != nullptr;
const size_t n_embd = llama_model_n_embd_inp(llama_get_model(ctx));
// positions continue from the memory when none are given
auto * mem = llama_get_memory(ctx);
std::vector<llama_pos> pos_next(llama_n_seq_max(ctx));
for (llama_seq_id s = 0; s < (llama_seq_id) pos_next.size(); ++s) {
pos_next[s] = llama_memory_seq_pos_max(mem, s) + 1;
}
if (!tokens.empty()) {
llama_batch_ext_set_output_logits(batch.get(), (int32_t) tokens.size() - 1, true);
for (int32_t i = 0; i < batch.n_tokens; ++i) {
const int32_t n_sid = batch.n_seq_id ? batch.n_seq_id[i] : 1;
const llama_seq_id seq_id = batch.seq_id ? batch.seq_id[i][0] : 0;
llama_pos pos[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
if (!batch.pos) {
pos[0] = pos_next[seq_id]++;
} else if (has_token) {
pos[0] = batch.pos[i];
} else {
// embedding batch: section-major layout pos[j*n_tokens + i]
for (int32_t j = 0; j < res.n_pos; ++j) {
pos[j] = batch.pos[j * batch.n_tokens + i];
}
}
const bool output = batch.logits ? batch.logits[i] != 0 : i == batch.n_tokens - 1;
const llama_embd embd = { has_embd ? batch.embd + (size_t) i * n_embd : nullptr, 1, n_embd };
int32_t idx;
if (has_token) {
idx = res.add(batch.token[i], pos[0], seq_id, output);
if (has_embd) {
res.set_embd(idx, embd);
}
} else {
idx = res.add_embd(embd, pos, seq_id, output);
}
for (int32_t s = 1; s < n_sid; ++s) {
llama_batch_ext_add_seq(res.get(), idx, batch.seq_id[i][s]);
}
}
return res;
}
common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
common_batch batch(ctx);
auto mem = llama_get_memory(ctx);
llama_pos pos = llama_memory_seq_pos_max(mem, 0) + 1; // -1 + 1 == 0 when the memory is empty
for (size_t i = 0; i < tokens.size(); ++i) {
const bool output = i == tokens.size() - 1;
batch.add(tokens[i], pos, 0, output);
pos++;
}
return batch;
@@ -2241,7 +2270,7 @@ bool common_prompt_batch_decode(
// memory, so we can't just remove the last token from the memory and replay the last token which
// is the reason for this logic.
llama_tokens prefix_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_tokens_before_last);
llama_batch_ext_ptr batch_prefix = common_batch_ext_get_one(ctx, prefix_tokens);
common_batch batch_prefix = common_batch_get_one(ctx, prefix_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_prefix.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
@@ -2251,10 +2280,8 @@ bool common_prompt_batch_decode(
llama_state_save_file(ctx, state_path.data(), all_tokens.data(), all_tokens.size());
COM_INF("saved session before last token to %s, n_new = %zu\n", state_path.data(), all_tokens.size());
llama_token last_token = all_tokens.back();
llama_batch_ext_ptr batch_last = common_batch_ext_get_one(ctx, { last_token });
llama_pos pos = n_past;
llama_batch_ext_set_pos(batch_last.get(), 0, &pos);
common_batch batch_last(ctx);
batch_last.add(all_tokens.back(), n_past, 0, true);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_last.get())) {
COM_ERR("%s", "failed to eval last token\n");
@@ -2263,7 +2290,7 @@ bool common_prompt_batch_decode(
n_past++;
} else {
llama_tokens new_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_new);
llama_batch_ext_ptr batch = common_batch_ext_get_one(ctx, new_tokens);
common_batch batch = common_batch_get_one(ctx, new_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
+68 -6
View File
@@ -8,6 +8,7 @@
#include "ggml.h"
#include "llama.h"
#include <array>
#include <list>
#include <set>
#include <sstream>
@@ -16,6 +17,7 @@
#include <vector>
#include <map>
#include <algorithm>
#include <filesystem>
#include <fstream>
#if defined(_WIN32) && !defined(_WIN32_WINNT)
@@ -808,7 +810,9 @@ static std::vector<T> string_split(const std::string & str, char delim) {
while (std::getline(str_stream, token, delim)) {
T value;
std::istringstream token_stream(token);
token_stream >> value;
if (!(token_stream >> value)) {
throw std::invalid_argument("invalid value: \"" + token + "\"");
}
values.push_back(value);
}
return values;
@@ -876,10 +880,21 @@ void string_process_escapes(std::string & input);
std::string string_from(bool value);
std::string string_from(const std::vector<int> & values);
std::string string_from(const struct llama_context * ctx, const std::vector<llama_token> & tokens);
std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch);
bool glob_match(const std::string & pattern, const std::string & str);
//
// Unicode utils
//
#ifdef _WIN32
std::wstring utf8_to_wstring(const std::string & str);
std::string wstring_to_utf8(const std::wstring & str);
#endif
// returns the path as a UTF-8 string, preserving its separators
std::string fs_path_to_utf8(const std::filesystem::path & path);
//
// Environment utils
//
@@ -889,16 +904,18 @@ bool glob_match(const std::string & pattern, const std::string & str);
std::string common_get_env(const std::string & name);
void common_set_env(const std::string & name, const std::string & value);
// reads a path from the environment, an unset variable gives an empty path
std::filesystem::path common_get_path_from_env(const std::string & name);
//
// Filesystem utils
//
bool fs_validate_filename(const std::string & filename, bool allow_subdirs = false);
bool fs_create_directory_with_parents(const std::string & path);
bool fs_is_directory(const std::string & path);
std::string fs_get_cache_directory();
std::string fs_get_cache_file(const std::string & filename);
std::filesystem::path fs_get_cache_directory();
std::filesystem::path fs_get_cache_file(const std::string & filename);
std::string fs_get_config_directory();
struct common_file_info {
@@ -1021,9 +1038,54 @@ void common_batch_add(
const std::vector<llama_seq_id> & seq_ids,
bool logits);
// wrapper around llama_batch_ext that provide getter functions for downstream code
struct common_batch {
struct token {
llama_token id;
std::array<llama_pos, GGML_MROPE_SECTIONS> pos; // only pos[0] is used for text tokens
llama_seq_id seq_id;
bool output;
llama_embd embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
};
std::vector<token> tokens; // mirror of the entries, tokens[i] describes batch index i
llama_batch_ext_ptr batch;
int32_t n_pos = 1; // positions per embedding entry, GGML_MROPE_SECTIONS for MROPE/IMROPE
common_batch() = default;
common_batch(struct llama_context * ctx);
llama_batch_ext * get() const { return batch.get(); }
// content type of the batch, all entries carry the same combination
bool has_token() const { return !tokens.empty() && tokens[0].id != LLAMA_TOKEN_NULL; }
bool has_embd () const { return !tokens.empty() && tokens[0].embd.data != nullptr; }
void clear();
// returns the batch index (>= 0), aborts if the entry cannot be added (batch full, invalid token or seq id)
int32_t add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output);
bool set_output(int32_t idx, bool value);
// 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);
// add an embedding-only entry (no token id), aborts like add() on failure
// pos points to n_pos positions
int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);
int32_t size() const { return (int32_t) tokens.size(); }
};
// create a single-sequence batch from a list of tokens
// last token always have output_logits set to true
llama_batch_ext_ptr common_batch_ext_get_one(struct llama_context * ctx, const llama_tokens & tokens);
common_batch common_batch_get_one(struct llama_context * ctx, const llama_tokens & tokens);
// convert a legacy llama_batch, applying its defaults: seq 0, positions continue from memory, last token is output
// the embd rows are read at the model input width
common_batch common_batch_from_llama_batch(struct llama_context * ctx, const llama_batch & batch);
// decodes a single batch of tokens for a prompt and manages session tokens
//
+2 -3
View File
@@ -1,4 +1,5 @@
#include "console.h"
#include "common.h"
#include "log.h"
#include <vector>
#include <iostream>
@@ -1053,9 +1054,7 @@ namespace console {
return false;
}
int size_needed = WideCharToMultiByte(CP_UTF8, 0, &wline[0], (int)wline.size(), NULL, 0, NULL, NULL);
line.resize(size_needed);
WideCharToMultiByte(CP_UTF8, 0, &wline[0], (int)wline.size(), &line[0], size_needed, NULL, NULL);
line = wstring_to_utf8(wline);
#else
if (!std::getline(std::cin, line)) {
// Input stream is bad or EOF received
+3 -3
View File
@@ -286,7 +286,7 @@ static int common_download_file_single_online(const std::string & url,
static const int max_attempts = 3;
static const int retry_delay_seconds = 2;
const bool file_exists = std::filesystem::exists(path);
const bool file_exists = std::filesystem::exists(std::filesystem::u8path(path));
if (file_exists && skip_etag) {
LOG_DBG("%s: using cached file: %s\n", __func__, path.c_str());
@@ -477,7 +477,7 @@ int common_download_file_single(const std::string & url,
return common_download_file_single_online(url, path, online_opts, skip_etag);
}
if (!std::filesystem::exists(path)) {
if (!std::filesystem::exists(std::filesystem::u8path(path))) {
LOG_ERR("%s: required file is not available in cache (offline mode): %s\n", __func__, path.c_str());
return -1;
}
@@ -943,7 +943,7 @@ std::string common_docker_resolve_model(const std::string & docker) {
std::string model_filename = repo;
std::replace(model_filename.begin(), model_filename.end(), '/', '_');
model_filename += "_" + tag + ".gguf";
std::string local_path = fs_get_cache_file(model_filename);
std::string local_path = fs_path_to_utf8(fs_get_cache_file(model_filename));
const std::string blob_url = url_prefix + "/blobs/" + gguf_digest;
common_download_opts opts;
+7 -7
View File
@@ -192,9 +192,9 @@ static void common_params_fit_impl(
uint32_t hp_nct = 0; // hparams.n_ctx_train
uint32_t hp_nex = 0; // hparams.n_expert
// size the context for all sequences, but keep minimums and alignment per KV stream
const uint32_t n_seq_max = std::max<uint32_t>(1, cparams->n_seq_max);
const uint32_t n_streams = cparams->kv_unified ? 1 : n_seq_max;
// with non-unified kv, we need to take into account n_streams
// for example, if memory can hold more than model's trained context size, we must extend the n_ctx to hold enough n_streams
const uint32_t n_streams = cparams->kv_unified ? 1 : std::max<uint32_t>(1, cparams->n_seq_max);
const bool n_ctx_auto = cparams->n_ctx == 0;
dmds_t dmds_extra; // memory of the extra model, laid out on the devices of the main model
@@ -264,15 +264,15 @@ static void common_params_fit_impl(
dmds_t dmds_full = common_get_device_memory_data_impl(path_model, mparams, cparams, devs, hp_ngl, hp_nct, hp_nex, log_level);
// saturate instead of overflowing, this also preserves the UINT32_MAX sentinel of n_ctx_min:
const uint32_t n_ctx_max = (uint32_t) std::min<uint64_t>(uint64_t(hp_nct) * n_seq_max, UINT32_MAX);
const uint32_t n_ctx_max = (uint32_t) std::min<uint64_t>(uint64_t(hp_nct) * n_streams, UINT32_MAX);
const uint32_t n_ctx_min_total = (uint32_t) std::min<uint64_t>(uint64_t(n_ctx_min) * n_streams, UINT32_MAX);
// llama_context would use only hp_nct in total for n_ctx == 0, resolve the context before measuring anything else:
if (n_ctx_auto) {
cparams->n_ctx = n_ctx_max;
if (n_seq_max > 1) {
LOG_TRC("%s: context size unset -> using %" PRIu32 " for %" PRIu32 " sequences:\n",
__func__, n_ctx_max, n_seq_max);
if (n_streams > 1) {
LOG_TRC("%s: context size unset and KV cache not unified -> using %" PRIu32 " for %" PRIu32 " sequences:\n",
__func__, n_ctx_max, n_streams);
dmds_full = common_get_device_memory_data_impl(path_model, mparams, cparams, devs, hp_ngl, hp_nct, hp_nex, log_level);
}
}
+28 -25
View File
@@ -44,8 +44,7 @@ static fs::path get_cache_directory() {
{HOME_DIR, fs::path(".cache") / "huggingface" / "hub"}
};
for (const auto & entry : entries) {
if (auto * p = std::getenv(entry.var); p && *p) {
fs::path base(p);
if (fs::path base = common_get_path_from_env(entry.var); !base.empty()) {
return entry.path.empty() ? base : base / entry.path;
}
}
@@ -63,12 +62,7 @@ static fs::path get_cache_directory() {
}
std::string get_cache_path() {
#if defined(__cpp_lib_char8_t)
const std::u8string u8str = get_cache_directory().u8string();
return std::string(reinterpret_cast<const char *>(u8str.data()), u8str.size());
#else
return get_cache_directory().u8string();
#endif
return fs_path_to_utf8(get_cache_directory());
}
static std::string folder_name_to_repo(const std::string & folder) {
@@ -179,7 +173,8 @@ static bool is_valid_subpath(const fs::path & path, const fs::path & subpath) {
}
static void safe_write_file(const fs::path & path, const std::string & data) {
fs::path path_tmp = path.string() + ".tmp";
fs::path path_tmp = path;
path_tmp += ".tmp";
if (path.has_parent_path()) {
fs::create_directories(path.parent_path());
@@ -196,7 +191,7 @@ static void safe_write_file(const fs::path & path, const std::string & data) {
}
if (file.fail() || ec) {
fs::remove(path_tmp, ec);
throw std::runtime_error("failed to write file: " + path.string());
throw std::runtime_error("failed to write file: " + fs_path_to_utf8(path));
}
}
@@ -246,6 +241,7 @@ static std::string get_repo_commit(const std::string & repo_id,
fs::path refs_path = get_repo_path(repo_id) / "refs";
std::string name;
std::string commit;
fs::path name_path;
for (const auto & branch : json["branches"]) {
if (!branch.is_object() ||
@@ -256,24 +252,28 @@ static std::string get_repo_commit(const std::string & repo_id,
std::string _name = branch["name"].get<std::string>();
std::string _commit = branch["targetCommit"].get<std::string>();
if (!is_valid_subpath(refs_path, _name)) {
LOG_WRN("%s: skip invalid branch: %s\n", __func__, _name.c_str());
continue;
}
if (!is_valid_commit(_commit)) {
LOG_WRN("%s: skip invalid commit: %s\n", __func__, _commit.c_str());
continue;
}
const fs::path candidate = fs::u8path(_name);
if (!is_valid_subpath(refs_path, candidate)) {
LOG_WRN("%s: skip invalid branch: %s\n", __func__, _name.c_str());
continue;
}
if (_name == "main") {
name = _name;
commit = _commit;
name_path = candidate;
break;
}
if (name.empty() || commit.empty()) {
name = _name;
commit = _commit;
name_path = candidate;
}
}
@@ -282,7 +282,7 @@ static std::string get_repo_commit(const std::string & repo_id,
return {};
}
safe_write_file(refs_path / name, commit);
safe_write_file(refs_path / name_path, commit);
return commit;
} catch (const common_json_error & e) {
@@ -331,7 +331,9 @@ hf_files get_repo_files(const std::string & repo_id,
file.repo_id = repo_id;
file.path = item["path"].get<std::string>();
if (!is_valid_subpath(commit_path, file.path)) {
const fs::path subpath = fs::u8path(file.path);
if (!is_valid_subpath(commit_path, subpath)) {
LOG_WRN("%s: skip invalid path: %s\n", __func__, file.path.c_str());
continue;
}
@@ -351,12 +353,12 @@ hf_files get_repo_files(const std::string & repo_id,
file.url = endpoint + repo_id + "/resolve/" + commit + "/" + file.path;
fs::path final_path = commit_path / file.path;
file.final_path = final_path.string();
fs::path final_path = commit_path / subpath;
file.final_path = fs_path_to_utf8(final_path);
if (!file.oid.empty() && !fs::exists(final_path)) {
fs::path local_path = blobs_path / file.oid;
file.local_path = local_path.string();
file.local_path = fs_path_to_utf8(local_path);
} else {
file.local_path = file.final_path;
}
@@ -423,7 +425,7 @@ hf_files get_cached_files(const std::string & repo_id) {
if (!fs::exists(snapshots_path)) {
continue;
}
std::string _repo_id = folder_name_to_repo(repo.path().filename().string());
std::string _repo_id = folder_name_to_repo(fs_path_to_utf8(repo.path().filename()));
if (!is_valid_repo_id(_repo_id)) {
continue;
@@ -446,8 +448,9 @@ hf_files get_cached_files(const std::string & repo_id) {
if (!path.empty()) {
hf_file file;
file.repo_id = _repo_id;
file.path = path.generic_string();
file.local_path = entry.path().string();
const auto generic_path = path.generic_u8string();
file.path = std::string(generic_path.begin(), generic_path.end());
file.local_path = fs_path_to_utf8(entry.path());
file.final_path = file.local_path;
files.push_back(std::move(file));
}
@@ -461,8 +464,8 @@ std::string finalize_file(const hf_file & file) {
static std::atomic<bool> symlinks_disabled{false};
std::error_code ec;
fs::path local_path(file.local_path);
fs::path final_path(file.final_path);
fs::path local_path = fs::u8path(file.local_path);
fs::path final_path = fs::u8path(file.final_path);
if (local_path == final_path || fs::exists(final_path, ec)) {
return file.final_path;
@@ -509,7 +512,7 @@ bool remove_cached_repo(const std::string & repo_id) {
std::error_code ec;
auto removed = fs::remove_all(repo_path, ec);
if (ec) {
LOG_ERR("%s: failed to remove repo cache %s: %s\n", __func__, repo_path.string().c_str(), ec.message().c_str());
LOG_ERR("%s: failed to remove repo cache %s: %s\n", __func__, fs_path_to_utf8(repo_path).c_str(), ec.message().c_str());
return false;
}
return removed > 0;
+9 -2
View File
@@ -429,8 +429,15 @@ private:
bool negate = false;
if (is_identifier("not")) { ++current; negate = true; }
auto test_id = parse_primary_expression();
// FIXME: tests can also be expressed like this: if x is eq 3
if (is(token::open_paren)) test_id = parse_call_expression(std::move(test_id));
if (is(token::open_paren)) {
test_id = parse_call_expression(std::move(test_id));
} else if (is(token::numeric_literal) || is(token::string_literal) || is(token::open_curly_bracket) || is(token::open_square_bracket) ||
(is(token::identifier) && !is_identifier("and") && !is_identifier("or") && !is_identifier("else"))) {
size_t call_pos = current;
statements args;
args.push_back(parse_unary_expression());
test_id = mk_stmt<call_expression>(call_pos, std::move(test_id), std::move(args));
}
operand = mk_stmt<test_expression>(start_pos, std::move(operand), negate, std::move(test_id));
}
return operand;
+61 -14
View File
@@ -348,6 +348,43 @@ static value default_value(const func_args & args) {
return no_value ? args.get_pos(1) : args.get_pos(0);
}
static value toobject(const func_args & args) {
auto out = mk_val<value_object>();
value iter = args.get_pos(0, mk_val<value_undefined>());
bool iter_first = false;
if (is_val<value_array>(iter)) {
iter_first = true;
for (const auto & it : iter->as_array()) {
if (is_val<value_array>(it) && it->as_array().size() == 2) {
auto tuple = it->as_array();
auto key = tuple[0];
auto val = tuple[1];
JJ_DEBUG("namespace/dict: adding key '%s'", key->as_string().str().c_str());
out->insert(key, val);
} else {
throw raised_exception("namespace/dict() iterable argument must consist of tuples, not " + it->type());
}
}
} else if (is_val<value_object>(iter)) {
iter_first = true;
for (const auto & pair : iter->as_ordered_object()) {
JJ_DEBUG("namespace/dict: adding key '%s'", pair.first->as_string().str().c_str());
out->insert(pair.first, pair.second);
}
}
for (const auto & arg : args.get_args()) {
if (is_val<value_kwarg>(arg)) {
auto kwarg = cast_val<value_kwarg>(arg);
JJ_DEBUG("namespace/dict: adding key '%s'", kwarg->key.c_str());
out->insert(kwarg->key, kwarg->val);
} else if (!iter_first) {
throw raised_exception("namespace/dict() arguments must be kwargs, dict and/or iterable of tuples, not " + arg->type());
}
iter_first = false;
}
return out;
}
const func_builtins & global_builtins() {
static const func_builtins builtins = {
{"raise_exception", [](const func_args & args) -> value {
@@ -355,18 +392,8 @@ const func_builtins & global_builtins() {
std::string msg = args.get_pos(0)->as_string().str();
throw raised_exception("Jinja Exception: " + msg);
}},
{"namespace", [](const func_args & args) -> value {
auto out = mk_val<value_object>();
for (const auto & arg : args.get_args()) {
if (!is_val<value_kwarg>(arg)) {
throw raised_exception("namespace() arguments must be kwargs");
}
auto kwarg = cast_val<value_kwarg>(arg);
JJ_DEBUG("namespace: adding key '%s'", kwarg->key.c_str());
out->insert(kwarg->key, kwarg->val);
}
return out;
}},
{"dict", toobject},
{"namespace", toobject},
{"strftime_now", [](const func_args & args) -> value {
args.ensure_vals<value_string>();
std::string format = args.get_pos(0)->as_string().str();
@@ -515,8 +542,28 @@ const func_builtins & global_builtins() {
}},
{"test_is_sameas", [](const func_args & args) -> value {
// Check if an object points to the same memory address as another object
(void)args;
throw not_implemented_exception("sameas test not implemented");
args.ensure_count(2);
auto a = args.get_pos(0);
auto b = args.get_pos(1);
bool res = false;
if (!is_val<value_undefined>(a) && !is_val<value_undefined>(b)) {
if (is_val<value_none>(a) && is_val<value_none>(b)) {
res = true;
} else if (is_val<value_bool>(a) && is_val<value_bool>(b)) {
if (a->as_bool() == b->as_bool()) {
res = true;
}
} else if (is_val<value_int>(a) && is_val<value_int>(b)) {
const int64_t x = a->as_int();
// Allow comparison within small-int cache range
if (x >= -5 && x <= 256 && x == b->as_int()) {
res = true;
}
} else if (a == b) {
res = true;
}
}
return mk_val<value_bool>(res);
}},
{"test_is_escaped", [](const func_args & args) -> value {
(void)args;
+14 -4
View File
@@ -43,9 +43,10 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
// Constrained grammar whenever tools are offered.
auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
auto has_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object();
// 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) {
auto start = p.rule("start", p.literal("<|start|>assistant"));
@@ -65,6 +66,15 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
auto final_msg = p.rule("final", recipient + p.literal("<|message|>") +
p.content(p.until_one_of({ "<|eot|>", "<|eom|>" })));
if (has_response_format) {
auto response_json = p.content(p.schema(p.json(), "response-format-schema", inputs.json_schema));
auto response_format = p.rule("response-format",
recipient + p.literal("<|message|>") +
((p.literal("```json") + p.space() + response_json + p.space() + p.literal("```")) | response_json));
return p.zero_or_more(start + analysis) + start + response_format;
}
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
auto string_value = p.ac(
p.tool_arg_string_value(p.until("</atem:parameter>")) + p.tool_arg_close(p.literal("</atem:parameter>")),
@@ -124,7 +134,7 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
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);
});
+1 -1
View File
@@ -214,7 +214,7 @@ struct common_sampler * common_sampler_init(
#ifdef LLAMA_USE_LLGUIDANCE
grmr = llama_sampler_init_llg(vocab, "lark", grammar_str.c_str());
#else
GGML_ABORT("llguidance (cmake -DLLAMA_LLGUIDANCE=ON) is not enabled");
throw std::runtime_error("failed to parse grammar: llguidance is not enabled");
#endif // LLAMA_USE_LLGUIDANCE
} else {
std::vector<std::string> trigger_patterns;
+183 -186
View File
@@ -165,7 +165,7 @@ struct common_speculative_impl {
virtual void begin(llama_seq_id seq_id, const llama_tokens & prompt) = 0;
virtual bool process(const llama_batch & batch) = 0;
virtual bool process(const common_batch & batch) = 0;
virtual void draft(common_speculative_draft_params_vec & dparams) = 0;
@@ -179,7 +179,11 @@ struct common_speculative_impl {
struct common_speculative_impl_draft_simple : public common_speculative_impl {
common_params_speculative_draft params;
llama_batch batch;
common_batch batch;
// zero row at the draft input width, stands in for target embeddings the draft cannot read
std::vector<float> zeros;
bool zeros_warned = false; // the substitution is reported once
std::vector<common_sampler_ptr> smpls;
@@ -194,6 +198,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
throw std::runtime_error("draft-simple requires a draft context");
}
zeros.assign(llama_model_n_embd_inp(llama_get_model(ctx_dft)), 0.0f);
SPC_TRC("%s", "adding speculative implementation 'draft-simple'\n");
SPC_TRC("- n_max=%d, n_min=%d, p_min=%f\n", this->params.n_max, this->params.n_min, this->params.p_min);
SPC_TRC("- gpu_layers=%d, cache_k=%s, cache_v=%s, ctx_tgt=%s, ctx_dft=%s, devices=[%s]\n",
@@ -204,7 +210,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
ctx_dft ? "yes" : "no",
common_speculative_get_devices_str(this->params.devices).c_str());
batch = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);
batch = common_batch(ctx_dft);
// TODO: optimize or pass from outside?
// {
@@ -249,21 +255,46 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
}
}
~common_speculative_impl_draft_simple() override {
llama_batch_free(batch);
}
void begin(llama_seq_id /*seq_id*/, const llama_tokens & /*prompt*/) override {
// noop
}
bool process(const llama_batch & batch) override {
bool process(const common_batch & batch_in) override {
auto * ctx_dft = params.ctx_dft;
llama_batch batch_dft = batch;
batch_dft.logits = nullptr;
// copy the entries to a batch owned by the draft context, only the last token is output
batch.clear();
const int32_t n_tokens = batch_in.size();
for (int32_t k = 0; k < n_tokens; ++k) {
const auto & t = batch_in.tokens[k];
const bool output = k == n_tokens - 1;
if (t.id != LLAMA_TOKEN_NULL) {
const int32_t idx = batch.add(t.id, t.pos[0], t.seq_id, output);
if (t.embd.data) {
batch.set_embd(idx, t.embd);
}
} else {
// mtmd input is projected by the target encoder, a draft with a different width cannot read it
// it gets zeros instead, keeping its positions contiguous
// ref: https://github.com/ggml-org/llama.cpp/pull/29385#discussion_r4124743243
const size_t n_embd = t.embd.n_rows * t.embd.n_embd;
const bool same_width = n_embd == zeros.size();
if (!same_width && !zeros_warned) {
SPC_WRN("target embeddings of size %zu do not fit the draft input width %zu, "
"the draft receives zero rows for them and drafts after multimodal input will be poor\n",
n_embd, zeros.size());
zeros_warned = true;
}
const llama_embd embd = same_width ? t.embd : llama_embd{ zeros.data(), 1, zeros.size() };
batch.add_embd(embd, t.pos.data(), t.seq_id, output);
}
}
const int ret = llama_decode(ctx_dft, batch_dft);
if (batch.size() == 0) {
return true;
}
const int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
SPC_ERR("failed to decode draft batch, ret = %d\n", ret);
@@ -277,7 +308,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
common_batch_clear(batch);
batch.clear();
// keep track of which sequences are still drafting
int n_drafting = 0;
@@ -294,12 +325,12 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
drafting[seq_id] = true;
common_sampler_reset(smpls[seq_id].get());
common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
batch.add(dp.id_last, dp.pos0, seq_id, true);
}
int ret = llama_decode(ctx_dft, batch);
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
SPC_ERR("llama_decode returned %d\n", ret);
SPC_ERR("llama_process returned %d\n", ret);
return;
}
@@ -308,7 +339,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
while (n_drafting > 0) {
int i_batch = 0;
common_batch_clear(batch);
batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (!drafting[seq_id]) {
@@ -353,17 +384,17 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
continue;
}
common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
batch.add(id, dp.pos0 + i + 1, seq_id, true);
}
if (batch.n_tokens == 0) {
if (batch.size() == 0) {
break;
}
// evaluate the drafted tokens on the draft model
ret = llama_decode(ctx_dft, batch);
ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
SPC_ERR("llama_process[%d] returned %d\n", i, ret);
break;
}
@@ -423,7 +454,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
// encoder+decoder on n_accepted+1 rows).
struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
common_params_speculative_draft params;
llama_batch batch;
common_batch batch; // decoder input, (token, g_embd) pairs
common_batch batch_enc; // encoder input, built from the extracted target features
std::vector<common_sampler_ptr> smpls;
@@ -477,11 +509,8 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
n_embd_enc = (int32_t) target_layer_ids_n * n_embd_tgt;
n_layer_tgt = llama_model_n_layer(model_tgt);
const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd_dec, /*n_seq_max=*/ 1);
// llama_batch_init allocates only one of token/embd; eagle3 decoder needs both.
// TODO: fix, how to call without malloc
batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
batch = common_batch(ctx_dft);
batch_enc = common_batch(ctx_dft);
smpls.resize(n_seq);
for (auto & s : smpls) {
@@ -543,12 +572,6 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
llama_sampler_free(backend_chains[seq_id]);
}
backend_chains.clear();
if (batch.token != nullptr) {
free(batch.token);
batch.token = nullptr;
}
llama_batch_free(batch);
}
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -567,16 +590,16 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
}
}
bool process(const llama_batch & batch_in) override {
if (batch_in.n_tokens <= 0) {
bool process(const common_batch & batch_in) override {
if (batch_in.size() <= 0) {
return true;
}
if (batch_in.token == nullptr || batch_in.embd != nullptr) {
if (!batch_in.has_token() || batch_in.has_embd()) {
return true;
}
const int32_t n_tokens = batch_in.n_tokens;
const int32_t n_tokens = batch_in.size();
// i_batch_beg[seq] / i_batch_end[seq]: inclusive batch indices of this seq's
// first/last token in batch_in. Assumes per-seq tokens are contiguous within
@@ -584,8 +607,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
std::vector<int32_t> i_batch_beg(n_seq, -1);
std::vector<int32_t> i_batch_end(n_seq, -1);
for (int k = 0; k < n_tokens; ++k) {
GGML_ASSERT(batch_in.n_seq_id[k] == 1);
const llama_seq_id seq_id = batch_in.seq_id[k][0];
const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
continue;
}
@@ -619,24 +641,23 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
g_embd_buf.resize((size_t) n_tokens * n_embd_dec);
// llama_encode() requires the full encoder batch to fit in n_ubatch.
// llama_process() requires the full encoder batch to fit in n_ubatch.
// Allow batch > ubatch: eagle3's per-token encoder can be chunked safely.
const int32_t n_ubatch_dft = (int32_t) llama_n_ubatch(ctx_dft);
for (int32_t i = 0; i < n_tokens; i += n_ubatch_dft) {
const int32_t n_chunk = std::min(n_ubatch_dft, n_tokens - i);
llama_batch enc_batch = {
/*.n_tokens =*/ n_chunk,
/*.token =*/ nullptr,
/*.embd =*/ features_buf.data() + (size_t) i * n_embd_enc,
/*.pos =*/ nullptr,
/*.n_seq_id =*/ nullptr,
/*.seq_id =*/ nullptr,
/*.logits =*/ nullptr,
};
const int32_t rc = llama_encode(ctx_dft, enc_batch);
// the per-token encoder does not use positions, generate placeholder ones from the memory state
batch_enc.clear();
llama_pos pos = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), 0) + 1;
for (int32_t j = 0; j < n_chunk; ++j) {
batch_enc.add_embd({ features_buf.data() + (size_t) (i + j) * n_embd_enc, 1, (size_t) n_embd_enc }, &pos, 0, true);
pos++;
}
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_ENCODE, batch_enc.get());
if (rc != 0) {
SPC_ERR("llama_encode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
rc, (int) n_chunk, (int) i);
return false;
}
@@ -664,7 +685,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
// deferred boundary, completed by the next process() or draft() call.
// (c) refresh deferred state — stash this ubatch's full g_embd into verify_g,
// update pending_g_last / pending_pos_last to the last row.
common_batch_clear(batch);
batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
const int32_t beg = i_batch_beg[seq_id];
@@ -679,36 +700,34 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
// 2) pending_pos_last + 1 == pos[beg]
// 3) pending_pos_last > dft_pos_max // TODO: is this check needed?
const llama_pos pending_pos = pending_pos_last[seq_id];
if (pending_pos >= 0 && pending_pos + 1 == batch_in.pos[beg]) {
if (pending_pos >= 0 && pending_pos + 1 == batch_in.tokens[beg].pos[0]) {
const llama_pos dft_pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id);
if (pending_pos > dft_pos_max) {
common_batch_add(batch, batch_in.token[beg], pending_pos, { seq_id }, /*logits=*/ false);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
pending_g_last[seq_id].data(), row_bytes);
const int32_t idx = batch.add(batch_in.tokens[beg].id, pending_pos, seq_id, /*output=*/ false);
batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
}
}
for (int32_t k = beg; k < end; ++k) {
common_batch_add(batch, batch_in.token[k + 1], batch_in.pos[k], { seq_id }, /*logits=*/ false);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
g_embd + (size_t) k * n_embd_dec, row_bytes);
const int32_t idx = batch.add(batch_in.tokens[k + 1].id, batch_in.tokens[k].pos[0], seq_id, /*output=*/ false);
batch.set_embd(idx, { g_embd + (size_t) k * n_embd_dec, 1, (size_t) n_embd_dec });
}
// refresh deferred state
const int32_t n_rows = end - beg + 1;
verify_pos_first[seq_id] = batch_in.pos[beg];
pending_pos_last[seq_id] = batch_in.pos[end];
verify_pos_first[seq_id] = batch_in.tokens[beg].pos[0];
pending_pos_last[seq_id] = batch_in.tokens[end].pos[0];
verify_g_rows[seq_id] = n_rows;
verify_g[seq_id].resize((size_t) n_rows * n_embd_dec, 0.0f);
std::memcpy(verify_g[seq_id].data(), g_embd + (size_t) beg * n_embd_dec, row_bytes * n_rows);
std::memcpy(pending_g_last[seq_id].data(), g_embd + (size_t) end * n_embd_dec, row_bytes);
}
if (batch.n_tokens > 0) {
const int32_t rc = llama_decode(ctx_dft, batch);
if (batch.size() > 0) {
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (rc != 0) {
SPC_ERR("llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
rc, (int) batch.n_tokens, (int) batch_in.pos[0]);
SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
rc, (int) batch.size(), (int) batch_in.tokens[0].pos[0]);
return false;
}
}
@@ -719,14 +738,12 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
common_batch_clear(batch);
batch.clear();
// keep track of which sequences are still drafting
int n_drafting = 0;
std::vector<bool> drafting(n_seq);
const size_t row_bytes = (size_t) n_embd_dec * sizeof(float);
// Complete the deferred boundary pair (dp.id_last, pending_g_last) at memory
// pos pending_pos_last. dp.id_last is target's freshest sample (= corrected
// token after verify, or first generated token after prefill), matching the
@@ -747,19 +764,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, pending_pos_last[seq_id], -1);
common_batch_add(batch, dp.id_last, pending_pos_last[seq_id], { seq_id }, true);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
pending_g_last[seq_id].data(),
row_bytes);
const int32_t idx = batch.add(dp.id_last, pending_pos_last[seq_id], seq_id, true);
batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
}
if (batch.n_tokens == 0) {
if (batch.size() == 0) {
return;
}
int ret = llama_decode(ctx_dft, batch);
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
SPC_ERR("llama_decode returned %d\n", ret);
SPC_ERR("llama_process returned %d\n", ret);
return;
}
@@ -768,7 +783,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
while (n_drafting > 0) {
int i_batch = 0;
common_batch_clear(batch);
batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (!drafting[seq_id]) {
@@ -814,17 +829,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
continue;
}
common_batch_add(batch, id, pending_pos_last[seq_id] + (i + 1), { seq_id }, true);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec, prenorm, row_bytes);
const int32_t idx = batch.add(id, pending_pos_last[seq_id] + (i + 1), seq_id, true);
batch.set_embd(idx, { prenorm, 1, (size_t) n_embd_dec });
}
if (batch.n_tokens == 0) {
if (batch.size() == 0) {
break;
}
ret = llama_decode(ctx_dft, batch);
ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
SPC_ERR("llama_process[%d] returned %d\n", i, ret);
break;
}
@@ -908,8 +923,10 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
struct common_speculative_impl_draft_dflash : public common_speculative_impl {
common_params_speculative_draft params;
llama_batch batch; // noise tokens
llama_batch batch_inject; // target features for KV cache injection
common_batch batch; // noise tokens
common_batch batch_inject; // target features for KV cache injection
std::vector<float> features_buf; // [n_chunk, n_embd_enc] gathered target features
std::vector<common_sampler_ptr> smpls;
@@ -1005,15 +1022,11 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
}
this->n_max = this->params.n_max;
batch = llama_batch_init(llama_n_batch(ctx_dft), 0, n_seq);
batch_inject = llama_batch_init(llama_n_ubatch(ctx_dft), n_embd_enc, n_seq);
batch = common_batch(ctx_dft);
batch_inject = common_batch(ctx_dft);
// embd batches on an M-RoPE draft need 4 position rows per token
// embd batches on an M-RoPE draft carry 4 position rows per token
is_mrope = llama_model_rope_type(model_dft) == LLAMA_ROPE_TYPE_MROPE;
if (is_mrope) {
free(batch_inject.pos);
batch_inject.pos = (llama_pos *) malloc(sizeof(llama_pos) * 4 * llama_n_batch(ctx_dft));
}
smpls.resize(n_seq);
for (auto & s : smpls) {
@@ -1062,9 +1075,6 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
llama_sampler_free(backend_chains[seq_id]);
}
backend_chains.clear();
llama_batch_free(batch);
llama_batch_free(batch_inject);
}
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -1085,8 +1095,8 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
}
}
bool process(const llama_batch & batch_in) override {
if (batch_in.n_tokens <= 0) {
bool process(const common_batch & batch_in) override {
if (batch_in.size() <= 0) {
return true;
}
@@ -1094,20 +1104,19 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
// produce the target-layer features used to seed the draft KV cache, so
// embeddings are injected too, except the pinned ones skipped below.
// TODO: revisit after https://github.com/ggml-org/llama.cpp/pull/24669 is merged
const bool has_tokens = batch_in.token != nullptr;
const bool has_embeddings = batch_in.embd != nullptr;
const bool has_tokens = batch_in.has_token();
const bool has_embeddings = batch_in.has_embd();
if (has_tokens == has_embeddings) {
return true;
}
const int32_t n_tokens = batch_in.n_tokens;
const int32_t n_tokens = batch_in.size();
// per-seq inclusive batch range (assumes each seq's tokens are contiguous in the batch)
std::vector<int32_t> i_batch_beg(n_seq, -1);
std::vector<int32_t> i_batch_end(n_seq, -1);
for (int32_t k = 0; k < n_tokens; ++k) {
GGML_ASSERT(batch_in.n_seq_id[k] == 1);
const llama_seq_id seq_id = batch_in.seq_id[k][0];
const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
continue;
}
@@ -1130,7 +1139,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
// an M-RoPE image pins all its rows to one position, so a windowed draft
// cache cannot free cells for it - skip it, the draft can jump over the gap
const bool pos_pinned = batch_in.pos[i_batch_beg[seq_id]] == batch_in.pos[i_batch_end[seq_id]];
const bool pos_pinned = batch_in.tokens[i_batch_beg[seq_id]].pos[0] == batch_in.tokens[i_batch_end[seq_id]].pos[0];
if (has_embeddings && n_rows > 1 && pos_pinned) {
continue;
}
@@ -1140,34 +1149,28 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
// gather target features per extract layer; the fused decode encodes and
// injects them into the K/V cache at the target positions
batch_inject.n_tokens = n_chunk;
features_buf.resize((size_t) n_chunk * n_embd_enc);
for (uint32_t k = 0; k < target_layer_ids_n; ++k) {
const float * layer = llama_get_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k]);
if (!layer) {
GGML_ABORT("DFlash: target layer %d input not extracted.", target_layer_ids[k]);
}
for (int32_t i = 0; i < n_chunk; ++i) {
float * dst = batch_inject.embd + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
float * dst = features_buf.data() + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
const float * src = layer + (size_t) (i_batch_beg[seq_id] + offset + i) * n_embd_tgt;
std::memcpy(dst, src, (size_t) n_embd_tgt * sizeof(float));
}
}
batch_inject.clear();
for (int32_t i = 0; i < n_chunk; ++i) {
const llama_pos p = batch_in.pos[i_batch_beg[seq_id] + offset + i];
batch_inject.pos[i] = p;
if (is_mrope) {
batch_inject.pos[1 * n_chunk + i] = p;
batch_inject.pos[2 * n_chunk + i] = p;
batch_inject.pos[3 * n_chunk + i] = 0;
}
batch_inject.n_seq_id[i] = 1;
batch_inject.seq_id[i][0] = seq_id;
batch_inject.logits[i] = false;
const llama_pos p = batch_in.tokens[i_batch_beg[seq_id] + offset + i].pos[0];
const llama_pos pos_arr[4] = { p, p, p, 0 };
batch_inject.add_embd({ features_buf.data() + (size_t) i * n_embd_enc, 1, (size_t) n_embd_enc }, pos_arr, seq_id, false);
}
const int32_t rc = llama_decode(ctx_dft, batch_inject);
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_inject.get());
if (rc != 0) {
LOG_ERR("%s: llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
LOG_ERR("%s: llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
__func__, rc, (int) n_chunk, (int) offset);
return false;
}
@@ -1180,7 +1183,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
common_batch_clear(batch);
batch.clear();
// build one batch holding every drafting sequence's noise block into a single decode)
// record where each block starts and its size
@@ -1200,21 +1203,21 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
const int32_t n_draft = params.n_max;
const int32_t n_block_tokens = n_draft + (is_dspark && sample_from_anchor ? 0 : 1);
i_block_beg[seq_id] = batch.n_tokens;
i_block_beg[seq_id] = batch.size();
n_block [seq_id] = n_block_tokens;
for (int32_t i = 0; i < n_block_tokens; ++i) {
common_batch_add(batch, i == 0 ? dp.id_last : mask_token_id, n + i, { seq_id }, !is_dflash2);
batch.add(i == 0 ? dp.id_last : mask_token_id, n + i, seq_id, !is_dflash2);
}
}
if (batch.n_tokens == 0) {
if (batch.size() == 0) {
return;
}
// decode all sequence's noise block in a single batch
int ret = llama_decode(ctx_dft, batch);
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
LOG_WRN("%s: llama_decode returned %d\n", __func__, ret);
LOG_WRN("%s: llama_process returned %d\n", __func__, ret);
return;
}
@@ -1328,7 +1331,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
struct common_speculative_impl_draft_mtp : public common_speculative_impl {
common_params_speculative_draft params; // reuses the draft-model params slot (ctx_tgt/ctx_dft)
llama_batch batch;
common_batch batch;
std::vector<common_sampler_ptr> smpls;
@@ -1384,11 +1387,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
ctx_dft ? "yes" : "no",
common_speculative_get_devices_str(this->params.devices).c_str());
const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd, /*n_seq_max=*/ 1);
// llama_batch_init allocates only one of token/embd; MTP needs both.
// TODO: fix, how to call without malloc
batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
batch = common_batch(ctx_dft);
smpls.resize(n_seq);
for (auto & s : smpls) {
@@ -1453,12 +1452,6 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
llama_sampler_free(backend_chains[seq_id]);
}
backend_chains.clear();
if (batch.token != nullptr) {
free(batch.token);
batch.token = nullptr;
}
llama_batch_free(batch);
}
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -1473,23 +1466,23 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
if (pos_max < N - 1 && !is_mem_shared) {
SPC_WRN("ctx_dft pos_max=%d < N-1=%d - "
"process() hook may not have run on every prefill ubatch "
"(need_embd / logits=1 on every prompt position?). "
"(need_embd / output flag on every prompt position?). "
"Drafts may degrade.\n",
(int) pos_max, N - 1);
}
}
bool process(const llama_batch & batch_in) override {
if (batch_in.n_tokens <= 0) {
bool process(const common_batch & batch_in) override {
if (batch_in.size() <= 0) {
return true;
}
// TODO: how to make it work with vision tokens?
if (batch_in.token == nullptr || batch_in.embd != nullptr) {
if (!batch_in.has_token() || batch_in.has_embd()) {
return true;
}
const int32_t n_tokens = batch_in.n_tokens;
const int32_t n_tokens = batch_in.size();
// remember the first and last batch index for each sequence
std::fill(i_batch_beg.begin(), i_batch_beg.end(), -1);
@@ -1497,9 +1490,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
for (int k = 0; k < n_tokens; ++k) {
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
GGML_ASSERT(batch_in.n_seq_id[k] == 1);
if (batch_in.seq_id[k][0] == seq_id) {
if (batch_in.tokens[k].seq_id == seq_id) {
i_batch_end[seq_id] = k;
if (i_batch_beg[seq_id] < 0) {
i_batch_beg[seq_id] = k;
@@ -1515,33 +1506,26 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
// if kv is shared with target (e.g Gemma4), then we can skip this catch-up decode
if (!is_mem_shared) {
common_batch_clear(batch);
batch.clear();
for (int k = 0; k < n_tokens; ++k) {
common_batch_add(batch, batch_in.token[k], batch_in.pos[k], { batch_in.seq_id[k][0] }, 0);
}
// shift the tgt embeddings to the right by one position
// pair each token with the tgt embedding shifted right by one position, and
// the first token of each sequence with the pending embedding from a previous run
// assumes that the tokens in the batch are sequential for each sequence
// i.e. we cannot have seq_id like this: [0, 0, 0, 1, 1, 0, 1, 1]
// ^--- this is a problem
// TODO:this is generally true, but would be nice to assert it
{
const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
std::memcpy(batch.embd + (size_t) 1 * n_embd, h_tgt, row_bytes * (n_tokens-1));
}
const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
// fill the pending embeddings from a previous run
auto set_h = [&](int idx, const float * h_row) {
std::memcpy(batch.embd + (size_t) idx * n_embd, h_row, row_bytes);
};
for (int k = 0; k < n_tokens; ++k) {
const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (i_batch_beg[seq_id] < 0) {
continue;
}
const int32_t idx = batch.add(batch_in.tokens[k].id, batch_in.tokens[k].pos[0], seq_id, false);
set_h(i_batch_beg[seq_id], pending_h[seq_id].data());
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 });
}
auto * mem_dft = llama_get_memory(ctx_dft);
@@ -1554,15 +1538,15 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
if (i_batch_beg[seq_id] < 0) {
continue;
}
llama_memory_seq_rm(mem_dft, seq_id, batch_in.pos[i_batch_beg[seq_id]], -1);
llama_memory_seq_rm(mem_dft, seq_id, batch_in.tokens[i_batch_beg[seq_id]].pos[0], -1);
}
llama_set_nextn_layer_offset(ctx_dft, head);
}
const int32_t rc = llama_decode(ctx_dft, batch);
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (rc != 0) {
SPC_ERR("llama_decode(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
head, (int) rc, (int) batch_in.pos[0]);
SPC_ERR("llama_process(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
head, (int) rc, (int) batch_in.tokens[0].pos[0]);
ok = false;
break;
}
@@ -1600,14 +1584,12 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
common_batch_clear(batch);
batch.clear();
// keep track of which sequences are still drafting
int n_drafting = 0;
std::vector<bool> drafting(n_seq);
const size_t row_bytes = (size_t) n_embd * sizeof(float);
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
auto & dp = dparams[seq_id];
@@ -1619,10 +1601,10 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
drafting[seq_id] = true;
common_sampler_reset(smpls[seq_id].get());
common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, pending_h[seq_id].data(), row_bytes);
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 });
i_last[seq_id] = batch.n_tokens - 1;
i_last[seq_id] = idx;
if (chain_heads) {
chain_h[seq_id].assign(pending_h[seq_id].begin(), pending_h[seq_id].end());
@@ -1648,16 +1630,16 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
llama_set_nextn_layer_offset(ctx_dft, i);
}
int ret = llama_decode(ctx_dft, batch);
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
SPC_ERR("llama_process[%d] returned %d\n", i, ret);
break;
}
// rebuild the batch for the next step: the growing-KV paths re-add only the
// new token (the KV already holds the prefix), while chained heads re-add the
// whole prefix at the next head. dropped sequences are simply not re-added.
common_batch_clear(batch);
batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (!drafting[seq_id]) {
@@ -1708,24 +1690,24 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
const int n_rows = (int) result.size() + 1; // id_last + tokens drafted so far
for (int t = 0; t < n_rows; ++t) {
const llama_token tok = (t == 0) ? dp.id_last : result[t - 1];
common_batch_add(batch, tok, dp.pos0 + t, { seq_id }, t == n_rows - 1);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd,
chain_h[seq_id].data() + (size_t) t * n_embd, row_bytes);
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 });
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
common_batch_add(batch, id, dp.pos0, { seq_id }, true);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
const int32_t idx = batch.add(id, dp.pos0, seq_id, true);
batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
i_last[seq_id] = idx;
} else {
common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
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 });
i_last[seq_id] = idx;
}
i_last[seq_id] = batch.n_tokens - 1;
}
if (batch.n_tokens == 0) {
if (batch.size() == 0) {
break;
}
@@ -1787,7 +1769,7 @@ struct common_speculative_impl_ngram_simple : public common_speculative_impl {
// noop
}
bool process(const llama_batch & /*batch*/) override {
bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -1835,7 +1817,7 @@ struct common_speculative_impl_ngram_map_k : public common_speculative_impl {
common_ngram_map_begin(config[seq_id], prompt);
}
bool process(const llama_batch & /*batch*/) override {
bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -1993,7 +1975,7 @@ struct common_speculative_impl_ngram_mod : public common_speculative_impl {
sinfo.n_draft_last = result.size();
}
bool process(const llama_batch & /*batch*/) override {
bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -2155,7 +2137,7 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
}
}
bool process(const llama_batch & /*batch*/) override {
bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -2181,6 +2163,9 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
struct common_speculative {
common_speculative_draft_params_vec dparams;
// the target context, used to convert legacy llama_batch inputs
llama_context * ctx_tgt = nullptr;
// list of implementations to use and their states
std::vector<std::unique_ptr<common_speculative_impl>> impls;
@@ -2726,6 +2711,7 @@ common_speculative * common_speculative_init(common_params_speculative & params,
common_speculative_ptr result(new common_speculative {
/* .dparams = */ common_speculative_draft_params_vec(n_seq),
/* .ctx_tgt = */ params.draft.ctx_tgt,
/* .impls = */ std::move(impls),
/* .impl_last = */ std::vector<common_speculative_impl *>(n_seq, nullptr),
/* .synth_probs = */ {},
@@ -2789,6 +2775,17 @@ void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, co
}
bool common_speculative_process(common_speculative * spec, const llama_batch & batch) {
if (spec == nullptr) {
return true;
}
// ngram-only setups have no target context, they do not read the batch anyway
const common_batch tmp = spec->ctx_tgt ? common_batch_from_llama_batch(spec->ctx_tgt, batch) : common_batch();
return common_speculative_process(spec, tmp);
}
bool common_speculative_process(common_speculative * spec, const common_batch & batch) {
bool result = true;
if (spec == nullptr) {
+3
View File
@@ -77,6 +77,9 @@ common_speculative_draft_params & common_speculative_get_draft_params(common_spe
void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, const llama_tokens & prompt);
// process the batch and update the internal state of the speculative context
bool common_speculative_process(common_speculative * spec, const common_batch & batch);
// legacy llama_batch input, converted with common_batch_from_llama_batch()
bool common_speculative_process(common_speculative * spec, const llama_batch & batch);
// generate drafts for the sequences specified with `common_speculative_get_draft_params`
+42
View File
@@ -170,6 +170,9 @@ class ModelBase:
self.dir_model_card = dir_model # overridden in convert_lora_to_gguf.py
self._is_nvfp4 = False
self._is_mxfp4 = False
self._nvfp4_global_algo: str | None = None # checkpoint-wide NVFP4 quant_algo
self._nvfp4_layer_algo: dict[str, str | None] = {} # per-layer quant_algo, keyed by HF module path
self._prec_a4: dict[str, bool] = {} # gguf tensor name -> can use 4-bit (A4) activations
self._fp8_as_q8 = fp8_as_q8
self._fp8_dequantized: set[str] = set()
@@ -664,6 +667,18 @@ class ModelBase:
if bias_types:
self._fusable_qkv_bias_layers.add(bid)
def _tag_prec_a4(self, hf_name: str, gguf_name: str) -> None:
# W4A16_NVFP4 should not use 4-bit activations
name = hf_name.removesuffix(".weight").removesuffix(".bias")
algo = self._nvfp4_global_algo
while name:
if name in self._nvfp4_layer_algo:
algo = self._nvfp4_layer_algo[name]
break
name = name.rpartition(".")[0]
if algo == "W4A16_NVFP4":
self._prec_a4[gguf_name] = False
def set_gguf_parameters(self):
raise NotImplementedError("set_gguf_parameters() must be implemented in subclasses")
@@ -837,6 +852,7 @@ class ModelBase:
raw, shape = self._nvfp4_pack(weight, scale)
logger.info(f"Repacked {new_name} with shape {shape} and quantization NVFP4")
self.gguf_writer.add_tensor(new_name, raw, raw_dtype=gguf.GGMLQuantizationType.NVFP4)
self._tag_prec_a4(name, new_name)
self._write_scale_tensor(new_name.replace(".weight", ".scale"), scale2)
self._write_scale_tensor(new_name.replace(".weight", ".input_scale"), input_scale)
@@ -929,6 +945,7 @@ class ModelBase:
new_name = self.map_tensor_name(merged_name)
logger.info(f"Repacked {new_name} with shape [{len(experts)}, {shape[0]}, {shape[1]}] and quantization NVFP4")
self.gguf_writer.add_tensor(new_name, merged, raw_dtype=gguf.GGMLQuantizationType.NVFP4)
self._tag_prec_a4(merged_name, new_name)
scales.sort(key=lambda x: x[0])
self._write_scales_tensor(new_name.replace(".weight", ".scale"), [s[1] for s in scales])
@@ -971,6 +988,9 @@ class ModelBase:
and bool(quant_groups)
and all(g.get("format") == "nvfp4-pack-quantized" for g in quant_groups.values() if isinstance(g, dict))
)
self._nvfp4_global_algo = quant_algo
if quant_algo != "NVFP4":
if nvfp4_compressed_tensors:
quant_algo = "NVFP4"
@@ -980,6 +1000,22 @@ class ModelBase:
self._is_nvfp4 = quant_algo in ("NVFP4", "W4A16_NVFP4")
self._is_mxfp4 = quant_method == "mxfp4"
# Per-tensor NVFP4 precision.
self._nvfp4_layer_algo = {}
if quant_layers:
# store all possible module paths and assert if a quantized layer is not in the model
modules: set[str] = set()
for name in self.model_tensors:
while name := name.rpartition(".")[0]:
modules.add(name)
for layer_name, entry in quant_layers.items():
if not isinstance(entry, dict):
continue
if titem := self.filter_tensors((layer_name, lambda: torch.empty(0))):
assert titem[0] in modules, f"quantized_layers entry {layer_name!r} is not in the model tensors"
self._nvfp4_layer_algo[titem[0]] = entry.get("quant_algo")
# NVFP4 weights are repacked and written directly to gguf_writer.
# This must run before dequant_model so NVFP4 tensors are removed
# from model_tensors, leaving only non-NVFP4 (e.g. FP8) for dequant.
@@ -1185,6 +1221,12 @@ class ModelBase:
logger.info("Set model quantization version")
self.gguf_writer.add_quantization_version(gguf.GGML_QUANT_VERSION)
if self._prec_a4:
names = sorted(self._prec_a4.keys())
values = [self._prec_a4[n] for n in names]
logger.info(f"Set prec_a4 metadata for {len(names)} tensor(s)")
self.gguf_writer.add_tensor_extra_prec_a4(names, values)
def write_vocab(self):
raise NotImplementedError("write_vocab() must be implemented in subclasses")
+1 -1
View File
@@ -341,7 +341,7 @@ class NomicBertModel(BertModel):
else:
raise ValueError(f"unrecognized parameters: n_positions={npos}, max_trained_positions={mtp}")
assert self.hparams["activation_function"] == "gelu" if self.is_moe else "swiglu"
assert self.hparams["activation_function"] == ("gelu" if self.is_moe else "swiglu")
# this doesn't do anything in the HF version
assert self.hparams["causal"] is False
+31 -3
View File
@@ -17,6 +17,7 @@ from .base import MmprojModel, ModelBase, TextModel, gguf, logger
@ModelBase.example("XiaomiMiMo/MiMo-V2.5")
class MimoV2Model(TextModel):
model_arch = gguf.MODEL_ARCH.MIMO2
supports_mtp_export = True
# MiMo V2-Flash, V2.5 and V2.5-Pro all ship 3 trained MTP layers under model.mtp.layers.{0,1,2}.
# The HF config does not expose the count, so it's hardcoded to match the count found in the safetensors.
@@ -25,6 +26,8 @@ class MimoV2Model(TextModel):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
if self.no_mtp:
self._n_nextn = 0
self.block_count = self.hparams["num_hidden_layers"] + self._n_nextn
self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count)
@@ -101,7 +104,7 @@ class MimoV2Model(TextModel):
qkv_overrides: dict[str, tuple[Callable, Callable, int]] = {}
qc = self.hparams.get("quantization_config")
if isinstance(qc, dict) and qc.get("quant_method") == "fp8":
pat = re.compile(r"^model\.layers\.(\d+)\.self_attn\.qkv_proj\.weight_scale_inv$")
pat = re.compile(r"^model\.(mtp\.)?layers\.(\d+)\.self_attn\.qkv_proj\.weight_scale_inv$")
for name in list(self.model_tensors.keys()):
m = pat.match(name)
if not m:
@@ -109,10 +112,13 @@ class MimoV2Model(TextModel):
weight_name = name.removesuffix("_scale_inv")
if weight_name not in self.model_tensors:
continue
bid = int(m.group(2))
if m.group(1) is not None:
bid += self.hparams["num_hidden_layers"]
qkv_overrides[weight_name] = (
self.model_tensors[weight_name],
self.model_tensors[name],
int(m.group(1)),
bid,
)
super().dequant_model()
@@ -165,7 +171,8 @@ class MimoV2Model(TextModel):
if v_scale is not None:
self.gguf_writer.add_attn_value_scale(float(v_scale))
self.gguf_writer.add_nextn_predict_layers(self._n_nextn)
if self._n_nextn > 0:
self.gguf_writer.add_nextn_predict_layers(self._n_nextn)
_MXFP4_EXPERT_RE = re.compile(
r"^model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.(gate|up|down)_proj\.weight$"
@@ -251,11 +258,32 @@ class MimoV2Model(TextModel):
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
name, gen = item
is_mtp = name.startswith("model.mtp.layers.")
if is_mtp and cls.no_mtp:
return None
if cls.mtp_only and not is_mtp and name not in (
"model.embed_tokens.weight", "model.norm.weight", "lm_head.weight",
):
return None
if "attention_sink" in name and not name.endswith(".weight"):
name += ".weight"
return super().filter_tensors((name, gen))
def prepare_metadata(self, vocab_only: bool):
from_dir = self.fname_out.is_dir()
super().prepare_metadata(vocab_only=vocab_only)
if not self.mtp_only or not from_dir:
return
output_type: str = self.ftype.name.partition("_")[2]
fname_default: str = gguf.naming_convention(
self.metadata.name, self.metadata.basename, self.metadata.finetune,
self.metadata.version, size_label=None, output_type=output_type, model_type=None)
self.fname_out = self.fname_out.parent / f"mtp-{fname_default}.gguf"
def modify_tensors(self, data_torch, name, bid):
# Remap MTP/NextN tensors to additional layer slots so the standard tensor map handles them.
# HF: model.mtp.layers.{i}.foo -> model.layers.{n_layer_text + i}.foo
-6
View File
@@ -424,12 +424,6 @@ class NemotronHModel(GraniteHybridModel):
special_vocab = gguf.SpecialVocab(self.dir_model, load_merges=True)
special_vocab.add_to_gguf(self.gguf_writer)
# The tokenizer _does_ add a BOS token (via post_processor type
# TemplateProcessing) but does not set add_bos_token to true in the
# config, so we need to explicitly override it here.
if not self.is_moe:
self.gguf_writer.add_add_bos_token(True)
_MTP_SPECIAL_RENAMES = {
"mtp.layers.0.enorm.weight": "model.layers.{bid}.enorm.weight",
"mtp.layers.0.hnorm.weight": "model.layers.{bid}.hnorm.weight",
+15
View File
@@ -154,6 +154,21 @@ class Plamo2Model(TextModel):
class Plamo3Model(TextModel):
model_arch = gguf.MODEL_ARCH.PLAMO3
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
# PLaMo-3 builds rope_parameters from flat config keys at runtime; mirror the YaRN settings for GGUF.
rope_scaling_factor = self.hparams.get("rope_scaling_factor", 1)
if rope_scaling_factor != 1 and "rope_type" not in self.rope_parameters:
self.rope_parameters.update({
"rope_type": "yarn",
"factor": float(rope_scaling_factor),
"original_max_position_embeddings": int(self.hparams["initial_context_length"]),
"beta_fast": 32.0,
"beta_slow": 1.0,
"truncate": False,
})
def set_vocab(self):
self._set_vocab_plamo()
+12 -4
View File
@@ -52,6 +52,10 @@ The packages for FP32 and FP16 would have different accuracy and performance on
## News
- 2026.09
- Update the CI build environment for oneAPI 2026.1 (unified oneAPI Toolkit). oneDNN is removed from the Deep Learning Essentials package in 2026.0, so the CI now uses the oneAPI Toolkit installer which still includes oneDNN.
- oneAPI 2026.1 improves the SYCL build performance: measured with the same code on Arc B570, prompt processing 1331 vs 434 t/s (3.1x) vs the 2025.3-based release build.
- 2026.04-05
- Optimize mul_mat by reorder feature for data type: Q4_K, Q5_K, Q6_K, Q8_0.
- Fused MoE.
@@ -257,7 +261,7 @@ Platform #0: Intel(R) OpenCL HD Graphics
`-- Device #0: Intel(R) Iris(R) Xe Graphics [0x9a49]
```
2. **Install Intel® oneAPI Base toolkit**
2. **Install Intel® oneAPI Toolkit**
SYCL backend depends on:
- Intel® oneAPI DPC++/C++ compiler/running-time.
@@ -267,11 +271,11 @@ SYCL backend depends on:
- **For Intel GPU**
All above are included in both **Intel® oneAPI Base toolkit** and **Intel® Deep Learning Essentials** packages.
With the 2026.0 release, the Intel® oneAPI Base toolkit and the HPC toolkit are combined into the **Intel® oneAPI Toolkit**, and **oneDNN is removed from the Intel® Deep Learning Essentials** package (oneDNN is distributed separately since then). The **Intel® oneAPI Toolkit** includes oneDNN until 2027.0.
It's recommended to install **Intel® Deep Learning Essentials** which only provides the necessary libraries with less size.
It's recommended to install the **Intel® oneAPI Toolkit**.
The **Intel® oneAPI Base toolkit** and **Intel® Deep Learning Essentials** can be obtained from the official [Intel® oneAPI Base Toolkit](https://www.intel.com/content/www/us/en/developer/tools/oneapi/base-toolkit.html) page.
The **Intel® oneAPI Toolkit** can be obtained from the official [Intel® oneAPI Toolkit](https://www.intel.com/content/www/us/en/developer/tools/oneapi/base-toolkit-download.html) page.
Please follow the instructions for downloading and installing the Toolkit for Linux, and preferably keep the default installation values unchanged, notably the installation path *(`/opt/intel/oneapi` by default)*.
@@ -281,6 +285,7 @@ Upon a successful installation, SYCL is enabled for the available Intel devices,
|Verified release|
|-|
|2026.1 |
|2025.3.3 |
|2025.2.1|
|2025.1|
@@ -811,6 +816,9 @@ 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_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.|
| GGML_SYCL_SPARSE_FA_MARGIN | [0,..] default:256 | Set the margin value for Sparse Flash-attention.|
| ZES_ENABLE_SYSMAN | 0 (default) or 1 | Support to get free memory of GPU by sycl::aspect::ext_intel_free_memory.<br>Recommended to use when --split-mode = layer |
| UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS | 0 (default) or 1 | Allow SYCL/Unified Runtime Level Zero device allocations larger than 4 GiB. llama.cpp's direct Level Zero allocation path requests the relaxed maximum-size limit itself when GGML_SYCL_ENABLE_LEVEL_ZERO=1. |
| GGML_SYCL_USM_SYSTEM | 0 (default) or 1 | Enable experimental support for [USM system allocations](https://github.khronos.org/SYCL_Reference/iface/usm_basic_concept.html#system-allocations) for large GPU buffers. This requires enough host memory for model weights and caches, an Intel Xe2+ GPU such as BMG or newer and supported on Linux only, with CONFIG_DRM_XE_GPUSVM enabled. |
@@ -77,6 +77,8 @@
{ "name": "arm64-android-snapdragon-debug" , "inherits": [ "base", "arm64-android-snapdragon", "debug" ] },
{ "name": "arm64-android-snapdragon-release", "inherits": [ "base", "arm64-android-snapdragon", "release" ] },
{ "name": "arm64-android-snapdragon-relwithdebinfo", "inherits": [ "arm64-android-snapdragon-release" ],
"cacheVariables": { "GGML_HEXAGON_HTP_BUILD_TYPE": "RelWithDebInfo" } },
{ "name": "arm64-windows-snapdragon-debug" , "inherits": [ "base", "arm64-windows-snapdragon", "debug" ] },
{ "name": "arm64-windows-snapdragon-release", "inherits": [ "base", "arm64-windows-snapdragon", "release" ] },
+17 -2
View File
@@ -181,6 +181,14 @@ cmake -B build -DGGML_CUDA=ON
cmake --build build --config Release
```
Note that this also builds the CPU backend by default. On Windows on ARM, MSVC's
support for the ARM NEON intrinsics used by the CPU backend may be incomplete, so
a CUDA build produced entirely with MSVC might have a slower CPU backend. If CPU
performance matters, try following the split build used in our release workflow
([.github/workflows/release.yml](../.github/workflows/release.yml)): the CPU backend
is built with clang (`cmake/arm64-windows-llvm.cmake`) and the CUDA backend with MSVC
(`cmake/arm64-windows-msvc-cuda.cmake`), and the artifacts are merged afterwards.
### Non-Native Builds
By default llama.cpp will be built for the hardware that is connected to the system at that time.
@@ -282,6 +290,13 @@ Consider setting `CUDA_SCALE_LAUNCH_QUEUES=4x`, which increases the CUDA command
Override default, speed-optimized compute types for cuBLAS matrix multiplications.
Legal values: `auto`, `f16`, `fp16`, `bf16`, `f32`, `fp32`.
#### GGML_CUDA_MMQ_PREC
Override the activation precision that the model requests for NVFP4 and MXFP4 matrix multiplications.
Currently supported values: `auto`, `q8`, `q4`.
NVFP4 and MXFP4 layers marked as W4A16 request 8-bit activations, so on Blackwell those layers run through the W4A8 path instead of the native W4A4 path. Set `q4` to keep the native W4A4 path for faster prompt processing at the cost of accuracy, or `q8` to use the W4A8 path for every layer, `auto` uses per-tensor prec metadata (this is the same behavior as when the environment variable is not set).
### Unified Memory
The environment variable `GGML_CUDA_ENABLE_UNIFIED_MEMORY=1` can be used to enable unified memory in Linux. This allows swapping to system RAM instead of crashing when the GPU VRAM is exhausted. In Windows this setting is available in the NVIDIA control panel as `System Memory Fallback`.
@@ -323,11 +338,11 @@ cmake --build build --config Release
By default, all supported compute capabilities are enabled. To customize this behavior, you can specify the `MUSA_ARCHITECTURES` option in the CMake command:
```bash
cmake -B build -DGGML_MUSA=ON -DMUSA_ARCHITECTURES="21"
cmake -B build -DGGML_MUSA=ON -DMUSA_ARCHITECTURES="31"
cmake --build build --config Release
```
This configuration enables only compute capability `2.1` (MTT S80) during compilation, which can help reduce compilation time.
This configuration enables only compute capability `3.1` (MTT S5000) during compilation, which can help reduce compilation time.
#### Compilation options
+1 -1
View File
@@ -123,7 +123,7 @@ You may want to pass in some different `ARGS`, depending on the MUSA environment
The defaults are:
- `MUSA_VERSION` set to `rc4.3.0`
- the base image is the MUSA 5.2.0 image from the Moore Threads registry
The resulting images, are essentially the same as the non-MUSA images:
+20 -16
View File
@@ -14,16 +14,16 @@ Legend:
| Operation | BLAS | CANN | CPU | CUDA | ET | HTP | MTL | OpenCL | SYCL | Vulkan | WebGPU | ZenDNN | zDNN |
|-----------|------|------|------|------|------|------|------|------|------|------|------|------|------|
| ABS | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ABS | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ACC | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ADD | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ADD1 | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| ADD_ID | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARANGE | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ARGMAX | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARGSORT | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARGMAX | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARGSORT | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| CEIL | ❌ | ❌ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| CLAMP | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| CLAMP | ❌ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| COL2IM_1D | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| CONCAT | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | 🟡 | ❌ | ❌ |
| CONT | ❌ | 🟡 | ✅ | ✅ | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | 🟡 | ❌ | ❌ |
@@ -55,7 +55,7 @@ Legend:
| GATED_LINEAR_ATTN | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| GEGLU | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GEGLU_ERF | ❌ | ✅ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GEGLU_QUICK | ❌ | ✅ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GEGLU_QUICK | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU_ERF | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU_QUICK | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
@@ -64,17 +64,21 @@ Legend:
| GROUP_NORM | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ |
| HARDSIGMOID | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| HARDSWISH | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| IM2COL | ❌ | ✅ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| IM2COL | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| IM2COL_3D | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| L2_NORM | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | 🟡 | ❌ | ❌ |
| LEAKY_RELU | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | 🟡 | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| LEAKY_RELU | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | 🟡 | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| LIGHTNING_INDEXER | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ |
| LOG | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| LOG | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MEAN | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | 🟡 | ✅ | ❌ | ❌ | ❌ |
| MUL | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MUL_MAT | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 |
| MUL_MAT_HADAMARD | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MUL_MAT_HADAMARD | ❌ | ❌ | ❌ | ❌ | ✅ | 🟡 | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MUL_MAT_ID | ❌ | 🟡 | ✅ | ✅ | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | 🟡 | ❌ |
| MUL_MAT_ID_W4A4 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| MUL_MAT_ID_W4A8 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| MUL_MAT_W4A4 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| MUL_MAT_W4A8 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| NEG | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| NORM | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | 🟡 | ❌ | ❌ |
| OPT_STEP_ADAMW | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
@@ -85,12 +89,12 @@ Legend:
| POOL_1D | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| POOL_2D | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| REGLU | ❌ | ✅ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| RELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| RELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| REPEAT | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | 🟡 | ❌ | ❌ |
| REPEAT_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| RMS_NORM | ❌ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| RMS_NORM_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROLL | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROLL | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROPE | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ROPE_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROUND | ❌ | ❌ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
@@ -108,20 +112,20 @@ Legend:
| SOFT_MAX | ❌ | 🟡 | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SOFT_MAX_BACK | ❌ | ❌ | 🟡 | 🟡 | ❌ | ❌ | ❌ | ❌ | 🟡 | ✅ | ❌ | ❌ | ❌ |
| SOLVE_TRI | ❌ | ❌ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ |
| SQR | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SQRT | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SQR | ❌ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SQRT | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SSM_CONV | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SSM_SCAN | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | 🟡 | 🟡 | ✅ | ❌ | ❌ |
| STEP | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| STEP | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SUB | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SUM | ❌ | 🟡 | ✅ | 🟡 | ❌ | ❌ | 🟡 | ❌ | 🟡 | 🟡 | 🟡 | ❌ | ❌ |
| SUM | ❌ | 🟡 | ✅ | 🟡 | ❌ | 🟡 | 🟡 | ❌ | 🟡 | 🟡 | 🟡 | ❌ | ❌ |
| SUM_ROWS | ❌ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | ❌ |
| SWIGLU | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SWIGLU_CLAMP | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ |
| SWIGLU_OAI | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| TANH | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| TIMESTEP_EMBEDDING | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ |
| TOP_K | ❌ | ❌ | ✅ | ❌ | ❌ | 🟡 | ✅ | ❌ | ✅ | 🟡 | ✅ | ❌ | ❌ |
| TOP_K | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | 🟡 | ✅ | ❌ | ❌ |
| TRI | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| TRUNC | ❌ | ❌ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| UPSCALE | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
+10632 -10041
View File
File diff suppressed because it is too large Load Diff
@@ -228,7 +228,6 @@ int main(int argc, char ** argv) {
common_batch_add(batch_tgt, draft[i], n_past + i, { seq_id }, true);
}
//LOG_DBG("target batch: %s\n", string_from(ctx_tgt, batch_tgt).c_str());
llama_decode(ctx_tgt, batch_tgt);
}
+1 -1
View File
@@ -5,7 +5,7 @@ project("ggml" C CXX ASM)
### GGML Version
set(GGML_VERSION_MAJOR 0)
set(GGML_VERSION_MINOR 25)
set(GGML_VERSION_PATCH 1)
set(GGML_VERSION_PATCH 3)
set(GGML_VERSION_BASE "${GGML_VERSION_MAJOR}.${GGML_VERSION_MINOR}.${GGML_VERSION_PATCH}")
list(APPEND CMAKE_MODULE_PATH "${CMAKE_CURRENT_SOURCE_DIR}/cmake/")
+1
View File
@@ -28,6 +28,7 @@ endfunction()
function(ggml_get_system_arch)
if (CMAKE_OSX_ARCHITECTURES STREQUAL "arm64" OR
CMAKE_GENERATOR_PLATFORM_LWR STREQUAL "arm64" OR
(CMAKE_GENERATOR_PLATFORM_LWR STREQUAL "arm64ec" AND MSVC AND NOT CMAKE_C_COMPILER_ID STREQUAL "Clang") OR
(NOT CMAKE_OSX_ARCHITECTURES AND NOT CMAKE_GENERATOR_PLATFORM_LWR AND
CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm.*|ARM64)$"))
set(GGML_SYSTEM_ARCH "ARM" PARENT_SCOPE)
+72 -33
View File
@@ -31,8 +31,6 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
ggml-cpu/ggml-cpu.cpp
ggml-cpu/repack.cpp
ggml-cpu/repack.h
ggml-cpu/iqp.cpp
ggml-cpu/iqp.h
ggml-cpu/hbm.cpp
ggml-cpu/hbm.h
ggml-cpu/quants.c
@@ -43,6 +41,10 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
ggml-cpu/amx/amx.h
ggml-cpu/amx/mmq.cpp
ggml-cpu/amx/mmq.h
ggml-cpu/tiled/tiled.cpp
ggml-cpu/tiled/tiled.h
ggml-cpu/tiled/tiled-kernel.cpp
ggml-cpu/tiled/tiled-kernel.h
ggml-cpu/ggml-cpu-impl.h
ggml-cpu/common.h
ggml-cpu/binary-ops.h
@@ -104,54 +106,79 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
ggml-cpu/arch/arm/repack.cpp
)
# MSVC proper (not clang-cl) has no -mcpu/-march for ARM: features can only be
# selected by probing the host at configure time and forwarding -D__ARM_FEATURE_*
if (MSVC AND NOT CMAKE_C_COMPILER_ID STREQUAL "Clang")
message(FATAL_ERROR "MSVC is not supported for ARM, use clang")
set(GGML_ARM_MSVC ON)
else()
check_cxx_compiler_flag(-mfp16-format=ieee GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E)
if (NOT "${GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E}" STREQUAL "")
list(APPEND ARCH_FLAGS -mfp16-format=ieee)
set(GGML_ARM_MSVC OFF)
endif()
if (GGML_ARM_MSVC AND NOT GGML_NATIVE)
message(FATAL_ERROR "MSVC on ARM requires GGML_NATIVE=ON, use clang for GGML_CPU_ARM_ARCH / GGML_CPU_ALL_VARIANTS")
else()
if (NOT GGML_ARM_MSVC)
check_cxx_compiler_flag(-mfp16-format=ieee GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E)
if (NOT "${GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E}" STREQUAL "")
list(APPEND ARCH_FLAGS -mfp16-format=ieee)
endif()
endif()
if (GGML_NATIVE)
# -mcpu=native does not always enable all the features in some compilers,
# so we check for them manually and enable them if available
execute_process(
COMMAND ${CMAKE_C_COMPILER} -mcpu=native -E -v -
INPUT_FILE "/dev/null"
OUTPUT_QUIET
ERROR_VARIABLE ARM_MCPU
RESULT_VARIABLE ARM_MCPU_RESULT
)
if (NOT ARM_MCPU_RESULT)
string(REGEX MATCH "-mcpu=[^ ']+" ARM_MCPU_FLAG "${ARM_MCPU}")
string(REGEX MATCH "-march=[^ ']+" ARM_MARCH_FLAG "${ARM_MCPU}")
# cl.exe cannot be queried this way: no /dev/null, and it rejects the flags.
# ARM_NATIVE_FLAG stays empty for MSVC and is never appended below.
if (NOT GGML_ARM_MSVC)
execute_process(
COMMAND ${CMAKE_C_COMPILER} -mcpu=native -E -v -
INPUT_FILE "/dev/null"
OUTPUT_QUIET
ERROR_VARIABLE ARM_MCPU
RESULT_VARIABLE ARM_MCPU_RESULT
)
if (NOT ARM_MCPU_RESULT)
string(REGEX MATCH "-mcpu=[^ ']+" ARM_MCPU_FLAG "${ARM_MCPU}")
string(REGEX MATCH "-march=[^ ']+" ARM_MARCH_FLAG "${ARM_MCPU}")
# on some old GCC we need to read -march=
if (ARM_MARCH_FLAG AND NOT "${ARM_MARCH_FLAG}" STREQUAL "-march=native")
set(ARM_NATIVE_FLAG "${ARM_MARCH_FLAG}")
elseif(ARM_MCPU_FLAG AND NOT "${ARM_MCPU_FLAG}" STREQUAL "-mcpu=native")
set(ARM_NATIVE_FLAG "${ARM_MCPU_FLAG}")
# on some old GCC we need to read -march=
if (ARM_MARCH_FLAG AND NOT "${ARM_MARCH_FLAG}" STREQUAL "-march=native")
set(ARM_NATIVE_FLAG "${ARM_MARCH_FLAG}")
elseif(ARM_MCPU_FLAG AND NOT "${ARM_MCPU_FLAG}" STREQUAL "-mcpu=native")
set(ARM_NATIVE_FLAG "${ARM_MCPU_FLAG}")
endif()
endif()
endif()
if ("${ARM_NATIVE_FLAG}" STREQUAL "")
set(ARM_NATIVE_FLAG -mcpu=native)
message(WARNING "ARM -march/-mcpu not found, -mcpu=native will be used")
else()
message(STATUS "ARM detected flags: ${ARM_NATIVE_FLAG}")
if ("${ARM_NATIVE_FLAG}" STREQUAL "")
set(ARM_NATIVE_FLAG -mcpu=native)
message(WARNING "ARM -march/-mcpu not found, -mcpu=native will be used")
else()
message(STATUS "ARM detected flags: ${ARM_NATIVE_FLAG}")
endif()
endif()
include(CheckCXXSourceRuns)
macro(check_arm_feature tag feature code)
set(CMAKE_REQUIRED_FLAGS_SAVE ${CMAKE_REQUIRED_FLAGS})
set(CMAKE_REQUIRED_FLAGS "${ARM_NATIVE_FLAG}+${tag}")
if (GGML_ARM_MSVC)
set(ARM_PROBE_FLAG "")
set(ARM_PROBE_NOFLAG "")
else()
set(ARM_PROBE_FLAG "${ARM_NATIVE_FLAG}+${tag}")
set(ARM_PROBE_NOFLAG "${ARM_NATIVE_FLAG}+no${tag}")
endif()
set(CMAKE_REQUIRED_FLAGS "${ARM_PROBE_FLAG}")
check_cxx_source_runs("${code}" GGML_MACHINE_SUPPORTS_${tag})
if (GGML_MACHINE_SUPPORTS_${tag})
set(ARM_NATIVE_FLAG_FIX "${ARM_NATIVE_FLAG_FIX}+${tag}")
if (GGML_ARM_MSVC)
# forward runtime probing decisions for MSVC
list(APPEND ARCH_FLAGS "-D__ARM_FEATURE_${feature}=1")
endif()
else()
set(CMAKE_REQUIRED_FLAGS "${ARM_NATIVE_FLAG}+no${tag}")
set(CMAKE_REQUIRED_FLAGS "${ARM_PROBE_NOFLAG}")
check_cxx_source_compiles("int main() { return 0; }" GGML_MACHINE_SUPPORTS_no${tag})
if (GGML_MACHINE_SUPPORTS_no${tag})
set(ARM_NATIVE_FLAG_FIX "${ARM_NATIVE_FLAG_FIX}+no${tag}")
@@ -161,12 +188,24 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
set(CMAKE_REQUIRED_FLAGS ${CMAKE_REQUIRED_FLAGS_SAVE})
endmacro()
check_arm_feature(dotprod DOTPROD "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vdotq_s32(_s, _a, _b); return 0; }")
check_arm_feature(i8mm MATMUL_INT8 "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vmmlaq_s32(_s, _a, _b); return 0; }")
check_arm_feature(sve SVE "#include <arm_sve.h>\nint main() { svfloat32_t _a, _b; volatile svfloat32_t _c = svadd_f32_z(svptrue_b8(), _a, _b); return 0; }")
if (GGML_ARM_MSVC)
# drop volatile
check_arm_feature(dotprod DOTPROD "#include <arm_neon.h>\nint main() { int8x16_t _a = vdupq_n_s8(0), _b = vdupq_n_s8(0); int32x4_t _s = vdupq_n_s32(0); _s = vdotq_s32(_s, _a, _b); return 0; }")
check_arm_feature(i8mm MATMUL_INT8 "#include <arm_neon.h>\nint main() { int8x16_t _a = vdupq_n_s8(0), _b = vdupq_n_s8(0); int32x4_t _s = vdupq_n_s32(0); _s = vmmlaq_s32(_s, _a, _b); return 0; }")
check_arm_feature(sve SVE "#include <arm_sve.h>\nint main() { const svbool_t pg = svptrue_b8(); const svint8_t a = svdup_n_s8(0), b = svdup_n_s8(0); svint32_t s = svdup_n_s32(0); s = svdot_s32(s, a, b); (void) pg; (void) s; return 0; }")
# FMA is mandatory on ARMv8-A but MSVC never defines __ARM_FEATURE_FMA,
# so probe it like the others and forward the define
check_arm_feature(fma FMA "#include <arm_neon.h>\nint main() { float32x4_t _a = vdupq_n_f32(0), _b = vdupq_n_f32(0), _c = vdupq_n_f32(0); _a = vfmaq_f32(_a, _b, _c); return 0; }")
else()
check_arm_feature(dotprod DOTPROD "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vdotq_s32(_s, _a, _b); return 0; }")
check_arm_feature(i8mm MATMUL_INT8 "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vmmlaq_s32(_s, _a, _b); return 0; }")
check_arm_feature(sve SVE "#include <arm_sve.h>\nint main() { svfloat32_t _a, _b; volatile svfloat32_t _c = svadd_f32_z(svptrue_b8(), _a, _b); return 0; }")
endif()
check_arm_feature(sme SME "#include <arm_sme.h>\n__arm_locally_streaming int main() { __asm__ volatile(\"smstart; smstop;\"); return 0; }")
list(APPEND ARCH_FLAGS "${ARM_NATIVE_FLAG}${ARM_NATIVE_FLAG_FIX}")
if (NOT GGML_ARM_MSVC)
list(APPEND ARCH_FLAGS "${ARM_NATIVE_FLAG}${ARM_NATIVE_FLAG_FIX}")
endif()
else()
if (GGML_CPU_ARM_ARCH)
list(APPEND ARCH_FLAGS -march=${GGML_CPU_ARM_ARCH})
+1 -1
View File
@@ -75,7 +75,7 @@
#define ggml_gemm_mxfp4_8x8_q8_0_generic ggml_gemm_mxfp4_8x8_q8_0
#define ggml_gemm_q8_0_4x4_q8_0_generic ggml_gemm_q8_0_4x4_q8_0
#define ggml_gemm_q8_0_4x8_q8_0_generic ggml_gemm_q8_0_4x8_q8_0
#elif defined(__aarch64__) || defined(__arm__) || defined(_M_ARM) || defined(_M_ARM64)
#elif defined(__aarch64__) || defined(__arm__) || defined(_M_ARM) || defined(_M_ARM64) || defined(_M_ARM64EC)
// repack.cpp
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
+23 -23
View File
@@ -875,23 +875,23 @@ void ggml_vec_dot_nvfp4_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const vo
const int8x8_t q8_3_lo = vld1_s8(y[2*ib+1].qs + 16);
const int8x8_t q8_3_hi = vld1_s8(y[2*ib+1].qs + 24);
const int32x4_t sumi = (int32x4_t){
const int32x4_t sumi = vld1q_s32(((const int32_t[4]) {
vaddvq_s32(ggml_nvfp4_dot8(q4_0_lo, q8_0_lo, q4_0_hi, q8_0_hi)),
vaddvq_s32(ggml_nvfp4_dot8(q4_1_lo, q8_1_lo, q4_1_hi, q8_1_hi)),
vaddvq_s32(ggml_nvfp4_dot8(q4_2_lo, q8_2_lo, q4_2_hi, q8_2_hi)),
vaddvq_s32(ggml_nvfp4_dot8(q4_3_lo, q8_3_lo, q4_3_hi, q8_3_hi)),
};
}));
#endif
const float dy0 = GGML_CPU_FP16_TO_FP32(y[2*ib].d);
const float dy1 = GGML_CPU_FP16_TO_FP32(y[2*ib+1].d);
const float32x4_t nvsc = {
const float32x4_t nvsc = vld1q_f32(((const float[4]) {
GGML_CPU_UE4M3_TO_FP32(x[ib].d[0]),
GGML_CPU_UE4M3_TO_FP32(x[ib].d[1]),
GGML_CPU_UE4M3_TO_FP32(x[ib].d[2]),
GGML_CPU_UE4M3_TO_FP32(x[ib].d[3])
};
const float32x4_t scales = vmulq_f32(nvsc, (float32x4_t){dy0, dy0, dy1, dy1});
}));
const float32x4_t scales = vmulq_f32(nvsc, vld1q_f32(((const float[4]) {dy0, dy0, dy1, dy1})));
acc = vfmaq_f32(acc, vcvtq_f32_s32(sumi), scales);
}
@@ -2591,7 +2591,7 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
memcpy(scales_mins, x0->scales, 12);
const uint32_t mins_0_3 = scales_mins[1] & kmask1;
const uint32_t mins_4_7 = ((scales_mins[2] >> 4) & kmask2) | (((scales_mins[1] >> 6) & kmask3) << 4);
const uint32x2_t mins = {mins_0_3, mins_4_7};
const uint32x2_t mins = vcreate_u32((uint64_t) mins_0_3 | ((uint64_t) mins_4_7 << 32));
x0_mins = vreinterpretq_s16_u16(vmovl_u8(vreinterpret_u8_u32(mins)));
uint32_t scales[2];
scales[0] = scales_mins[0] & kmask1; // scales 0~3
@@ -2603,7 +2603,7 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
memcpy(scales_mins, x1->scales, 12);
const uint32_t mins_0_3 = scales_mins[1] & kmask1;
const uint32_t mins_4_7 = ((scales_mins[2] >> 4) & kmask2) | (((scales_mins[1] >> 6) & kmask3) << 4);
const uint32x2_t mins = {mins_0_3, mins_4_7};
const uint32x2_t mins = vcreate_u32((uint64_t) mins_0_3 | ((uint64_t) mins_4_7 << 32));
x1_mins = vreinterpretq_s16_u16(vmovl_u8(vreinterpret_u8_u32(mins)));
uint32_t scales[2];
scales[0] = scales_mins[0] & kmask1; // scales 0~3
@@ -2611,7 +2611,7 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
memcpy(x1_scales, scales, 8);
}
int32x4_t visum = {0};
int32x4_t visum = vdupq_n_s32(0);
// process 64 data points per iteration, totally 256 data points
for (int j = 0; j < QK_K / 64; ++j, qx0 += 32, qx1 += 32, qy0 += 64, qy1 += 64) {
@@ -2637,14 +2637,14 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
// process 32 data points (share same block scale) per iteration
for (int k = 0; k < 2; ++k) {
const int blk = j * 2 + k;
const int32x4_t block_scale = {
const int32x4_t block_scale = vld1q_s32(((const int32_t[4]) {
x0_scales[blk],
x0_scales[blk],
x1_scales[blk],
x1_scales[blk],
};
}));
int32x4_t vr = {0};
int32x4_t vr = vdupq_n_s32(0);
for (int l = 0; l < 2; ++l) {
const int idx = k * 2 + l;
const int64x2_t vx0_s64 = vreinterpretq_s64_s8(vx0[idx]);
@@ -2678,20 +2678,22 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
vmull_s16(vget_high_s16(y0_sums), vget_high_s16(x1_mins))));
bias[3] = vaddvq_s32(vaddq_s32(vmull_s16(vget_low_s16(y1_sums), vget_low_s16(x1_mins)),
vmull_s16(vget_high_s16(y1_sums), vget_high_s16(x1_mins))));
const float32x4_t dmins = {
// note: the parentheses around the compound literal are required, vld1q_f32() is a
// function-like macro on MSVC and would otherwise see four separate arguments
const float32x4_t dmins = vld1q_f32(((const float[4]) {
GGML_CPU_FP16_TO_FP32(x0->dmin) * y0->d,
GGML_CPU_FP16_TO_FP32(x0->dmin) * y1->d,
GGML_CPU_FP16_TO_FP32(x1->dmin) * y0->d,
GGML_CPU_FP16_TO_FP32(x1->dmin) * y1->d,
};
}));
vfsum = vmlsq_f32(vfsum, vcvtq_f32_s32(vld1q_s32(bias)), dmins);
const float32x4_t superblock_scale = {
const float32x4_t superblock_scale = vld1q_f32(((const float[4]) {
GGML_CPU_FP16_TO_FP32(x0->d) * y0->d,
GGML_CPU_FP16_TO_FP32(x0->d) * y1->d,
GGML_CPU_FP16_TO_FP32(x1->d) * y0->d,
GGML_CPU_FP16_TO_FP32(x1->d) * y1->d,
};
}));
vfsum = vmlaq_f32(vfsum, vcvtq_f32_s32(visum), superblock_scale);
}
}
@@ -3261,12 +3263,8 @@ void ggml_vec_dot_q6_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
qy0 += 16;
qy1 += 16;
const int32x4_t block_scale = {
x0->scales[blk],
x0->scales[blk],
x1->scales[blk],
x1->scales[blk],
};
const int32x4_t block_scale =
vcombine_s32(vdup_n_s32(x0->scales[blk]), vdup_n_s32(x1->scales[blk]));
// calculate four results at once with outer product
const int8x16_t vx_l = vreinterpretq_s8_s64(vzip1q_s64(vreinterpretq_s64_s8(vx0[k]), vreinterpretq_s64_s8(vx1[k])));
@@ -3319,12 +3317,14 @@ void ggml_vec_dot_q6_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
const int32x4_t vibias = vmulq_n_s32(vld1q_s32(bias), 32);
const float32x4_t superblock_scale = {
// note: the parentheses around the compound literal are required, vld1q_f32() is a
// function-like macro on MSVC and would otherwise see four separate arguments
const float32x4_t superblock_scale = vld1q_f32(((const float[4]) {
GGML_CPU_FP16_TO_FP32(x0->d) * y0->d,
GGML_CPU_FP16_TO_FP32(x0->d) * y1->d,
GGML_CPU_FP16_TO_FP32(x1->d) * y0->d,
GGML_CPU_FP16_TO_FP32(x1->d) * y1->d,
};
}));
visum = vsubq_s32(visum, vibias);
vfsum = vmlaq_f32(vfsum, vcvtq_f32_s32(visum), superblock_scale);
+18 -32
View File
@@ -4,7 +4,6 @@
#include "ggml-backend-impl.h"
#include "ggml-backend.h"
#include "traits.h"
#include "iqp.h"
#include "ggml-cpu-impl.h"
#include "ggml-impl.h"
#include "quants.h"
@@ -15,6 +14,7 @@
#include "ops.h"
#include "ggml.h"
#include "common.h"
#include "tiled/tiled.h"
#if defined(_MSC_VER) || defined(__MINGW32__)
#include <malloc.h> // using malloc.h with MSC/MINGW
@@ -459,7 +459,7 @@ typedef pthread_mutex_t ggml_mutex_t;
#define ggml_lock_init(x) UNUSED(x)
#define ggml_lock_destroy(x) UNUSED(x)
#if defined(__x86_64__) || (defined(_MSC_VER) && defined(_M_AMD64))
#if defined(__x86_64__) || (defined(_MSC_VER) && defined(_M_AMD64) && !defined(_M_ARM64EC))
#define ggml_lock_lock(x) _mm_pause()
#else
#define ggml_lock_lock(x) UNUSED(x)
@@ -1265,6 +1265,11 @@ void ggml_compute_forward_mul_mat(
return;
}
// If tiled is supported, it will execute the full op here and we return
if (ggml_compute_forward_mul_mat_tiled(params, dst)) {
return;
}
GGML_TENSOR_BINARY_OP_LOCALS
const int ith = params->ith;
@@ -1372,13 +1377,6 @@ UseGgmlGemm1:;
ggml_barrier(params->threadpool);
// IQ panel gemm (see iqp.h) - must come after the barrier above, it consumes the q8_K rows
// of src1 from the work buffer
if (ggml_cpu_iqp_supports_mul_mat(dst) && !params->use_ref) {
ggml_compute_forward_mul_mat_iqp(params, dst);
return;
}
#if GGML_USE_LLAMAFILE
if (src1->type != vec_dot_type) {
const void* wdata = (src1->type == vec_dot_type) ? src1->data : params->wdata;
@@ -1596,15 +1594,9 @@ static void ggml_compute_forward_mul_mat_id(
char (*atomic_current_chunk)[CACHE_LINE_SIZE] = // [n_as]
incr_ptr_aligned(&wdata_cur, CACHE_LINE_SIZE * n_as, CACHE_LINE_SIZE);
// IQ panel gemm (see iqp.h); per expert eligibility is decided below, but the work buffer is
// reserved for the whole node (ggml_graph_plan sizes it without params, use_ref only skips the dispatch)
const bool iqp = ggml_cpu_iqp_supports_mul_mat_id(dst) && !params->use_ref;
char * iqp_panels = NULL;
if (iqp) {
iqp_panels = incr_ptr_aligned(&wdata_cur, nth * ggml_cpu_iqp_scratch_size(dst), 64);
}
// Tiled matmul (see tiled.h); per-thread work buffers, 0 bytes when disabled. The
// reservation is unconditional, the per expert eligibility is decided at dispatch time
char * tiled_scratch = incr_ptr_aligned(&wdata_cur, ggml_tiled_wdata_size(nth, dst), 64);
GGML_ASSERT(params->wsize >= (size_t)((char *) wdata_cur - (char *) params->wdata));
@@ -1677,10 +1669,8 @@ static void ggml_compute_forward_mul_mat_id(
continue;
}
if (iqp && ggml_cpu_iqp_mul_mat_id_min_batch(cne1)) {
ggml_compute_forward_mul_mat_id_iqp(params, dst, cur_a, cne1, (const int32_t *) &MMID_MATRIX_ROW(cur_a, 0),
iqp_panels);
// tiled takes over if profitable for this expert (see tiled.h)
if (ggml_compute_forward_mul_mat_id_tiled(params, dst, cur_a, cne1, (const int32_t *) &MMID_MATRIX_ROW(cur_a, 0), tiled_scratch)) {
continue;
}
@@ -2887,15 +2877,12 @@ struct ggml_cplan ggml_graph_plan(
case GGML_OP_MUL_MAT:
{
const enum ggml_type vec_dot_type = type_traits_cpu[node->src[0]->type].vec_dot_type;
if (node->src[1]->type != vec_dot_type) {
cur = ggml_row_size(vec_dot_type, ggml_nelements(node->src[1]));
}
// the IQ panel path needs one scratch panel per thread past the q8_K rows
if (ggml_cpu_iqp_supports_mul_mat(node)) {
cur = GGML_PAD(cur, 64) + n_tasks * ggml_cpu_iqp_scratch_size(node);
}
// Workspace for tiled (see tiled.h)
cur = GGML_PAD(cur, 64);
cur += ggml_tiled_wdata_size(n_tasks, node);
} break;
case GGML_OP_MUL_MAT_ID:
{
@@ -2915,10 +2902,9 @@ struct ggml_cplan ggml_graph_plan(
cur += n_as*ids->ne[0]*ids->ne[1]*sizeof(struct mmid_row_mapping) + sizeof(int64_t);
// atomic_current_chunk
cur += CACHE_LINE_SIZE*n_as + CACHE_LINE_SIZE;
// the IQ panel path needs one scratch panel per thread on top of that
if (ggml_cpu_iqp_supports_mul_mat_id(node)) {
cur += n_tasks * ggml_cpu_iqp_scratch_size(node) + 64;
}
// Workspace for tiled (see tiled.h)
cur = GGML_PAD(cur, 64);
cur += ggml_tiled_wdata_size(n_tasks, node);
} break;
case GGML_OP_OUT_PROD:
{
File diff suppressed because it is too large Load Diff
-39
View File
@@ -1,39 +0,0 @@
#pragma once
#include "ggml-cpu-impl.h"
#include "ggml.h"
// GGML internal header
// batched mul_mat path for the grid based IQ types: decode 8 src0 rows at a time into per thread scratch
// (block_iqp_x8, see iqp.cpp) and run an integer gemm over them against all src1 columns
#ifdef __cplusplus
extern "C" {
#endif
// whether cne1 rows of src1 are enough for the decode to pay for itself, per expert, for MUL_MAT_ID
bool ggml_cpu_iqp_mul_mat_id_min_batch(int64_t cne1);
bool ggml_cpu_iqp_supports_mul_mat(const struct ggml_tensor * dst);
// node level test only - per expert eligibility is decided with ggml_cpu_iqp_mul_mat_id_min_batch
bool ggml_cpu_iqp_supports_mul_mat_id(const struct ggml_tensor * dst);
// per thread panel scratch bytes, padded
size_t ggml_cpu_iqp_scratch_size(const struct ggml_tensor * dst);
// must be called after src1 has been converted to q8_K into params->wdata and the threads have synchronized on it
void ggml_compute_forward_mul_mat_iqp(const struct ggml_compute_params * params, struct ggml_tensor * dst);
// one expert: expert_rows points at its row of the matrix_rows table of (i1, i2) int32 pairs, panels at the base of the per thread panel scratches
void ggml_compute_forward_mul_mat_id_iqp(const struct ggml_compute_params * params,
struct ggml_tensor * dst,
int64_t cur_a,
int64_t cne1,
const int32_t * expert_rows,
void * panels);
#ifdef __cplusplus
}
#endif
+1 -1
View File
@@ -228,7 +228,7 @@ template <> inline vfloat32m8_t madd(vbfloat16m4_t a, vbfloat16m4_t b, vfloat32m
////////////////////////////////////////////////////////////////////////////////////////////////////
// VECTORIZED HORIZONTAL SUM
#if defined(__ARM_NEON)
#if defined(__ARM_NEON) || defined(_M_ARM64) || defined(_M_ARM64EC)
inline float hsum(float32x4_t x) {
return vaddvq_f32(x);
}
+6 -6
View File
@@ -9037,6 +9037,11 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
simd_gemm(KQ, (const float *)Q_q, K_f32, Q_TILE_SZ, DK, KV_TILE_SZ);
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, scale);
if (logit_softcap != 0.0f) {
ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
}
// Set padded KQ entries to -inf so softmax gives them zero weight
if (kv_tile < KV_TILE_SZ) {
for (int tq = 0; tq < Q_TILE_SZ; tq++) {
@@ -9046,11 +9051,6 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
}
}
if (logit_softcap != 0.0f) {
ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
}
if (mask) {
ggml_vec_add_f32(tile_rows * KV_TILE_SZ, KQ, KQ, mask32);
}
@@ -9320,7 +9320,7 @@ static void ggml_compute_forward_flash_attn_ext_f16(
kv_is_f32_or_f16 &&
k->type == v->type &&
neq1 >= Q_TILE_SZ);
#ifdef GGML_SIMD
#if defined(GGML_SIMD) && !defined(__x86_64__) && !defined(_M_X64)
#if defined(__ARM_FEATURE_SVE)
const int64_t f32_epr = svcntw();
#else
+54 -14
View File
@@ -56,6 +56,56 @@ static inline void simd_gemm_ukernel(
}
}
template <int RM>
static inline void simd_gemm_ukernel_tail(
float * GGML_RESTRICT C,
const float * GGML_RESTRICT A,
const float * GGML_RESTRICT B,
int K, int N, int cols)
{
#if defined(__AVX512F__)
const __mmask16 mask = (1u << cols) - 1;
__m512 acc[RM];
for (int64_t i = 0; i < RM; i++) {
acc[i] = _mm512_maskz_loadu_ps(mask, C + i * N);
}
for (int64_t kk = 0; kk < K; kk++) {
const __m512 b = _mm512_maskz_loadu_ps(mask, B + kk * N);
for (int64_t i = 0; i < RM; i++) {
acc[i] = _mm512_mask3_fmadd_ps(_mm512_set1_ps(A[i * K + kk]), b, acc[i], mask);
}
}
for (int64_t i = 0; i < RM; i++) {
_mm512_mask_storeu_ps(C + i * N, mask, acc[i]);
}
#elif defined(__AVX2__)
const __m256i mask = _mm256_cmpgt_epi32(_mm256_set1_epi32(cols), _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7));
__m256 acc[RM];
for (int64_t i = 0; i < RM; i++) {
acc[i] = _mm256_maskload_ps(C + i * N, mask);
}
for (int64_t kk = 0; kk < K; kk++) {
const __m256 b = _mm256_maskload_ps(B + kk * N, mask);
for (int64_t i = 0; i < RM; i++) {
acc[i] = GGML_F32_VEC_FMA(acc[i], b, _mm256_set1_ps(A[i * K + kk]));
}
}
for (int64_t i = 0; i < RM; i++) {
_mm256_maskstore_ps(C + i * N, mask, acc[i]);
}
#else
for (int64_t j = 0; j < cols; j++) {
for (int64_t i = 0; i < RM; i++) {
float a = C[i * N + j];
for (int64_t kk = 0; kk < K; kk++) {
a += A[i * K + kk] * B[kk * N + j];
}
C[i * N + j] = a;
}
}
#endif
}
// C[M x N] += A[M x K] * B[K x N]
static void simd_gemm(
float * GGML_RESTRICT C,
@@ -74,14 +124,8 @@ static void simd_gemm(
for (; jj + KN <= N; jj += KN) {
simd_gemm_ukernel<GEMM_RM, 1>(C + jj, A, B + jj, K, N);
}
for (; jj < N; jj++) {
for (int64_t i = 0; i < GEMM_RM; i++) {
float a = C[i * N + jj];
for (int64_t kk = 0; kk < K; kk++) {
a += A[i * K + kk] * B[kk * N + jj];
}
C[i * N + jj] = a;
}
if (jj < N) {
simd_gemm_ukernel_tail<GEMM_RM>(C + jj, A, B + jj, K, N, N - jj);
}
A += GEMM_RM * K;
@@ -97,12 +141,8 @@ static void simd_gemm(
for (; jj + KN <= N; jj += KN) {
simd_gemm_ukernel<1, 1>(C + jj, A, B + jj, K, N);
}
for (; jj < N; jj++) {
float a = C[jj];
for (int64_t kk = 0; kk < K; kk++) {
a += A[kk] * B[kk * N + jj];
}
C[jj] = a;
if (jj < N) {
simd_gemm_ukernel_tail<1>(C + jj, A, B + jj, K, N, N - jj);
}
A += K;
+719
View File
@@ -0,0 +1,719 @@
// mulmat microtile kernels
#include "tiled-kernel.h"
#include "ggml.h"
#include <cstring>
#if defined(__AVX512VNNI__) || defined(__AVX2__) || defined(__AVX__)
#include <immintrin.h>
#endif
// Reference implementation, slower than existing vec_dot approach
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
static void tiled_run_microtile_scalar(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
constexpr int NB = TILED_TILE_K / SUBBLK; // subblocks per 256-K block
constexpr int NS = SUBBLK / 16; // per-16 bsums per subblock
// num_k/slab: the narrow path holds num_k slabs at row stride num_k*256 (num_k=1, slab=0 = standard)
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
// version serves all of it.
const int nkr = (NK > 0) ? NK : num_k;
const int seff = (NK > 0) ? 0 : slab;
const int qk_stride = nkr * TILED_TILE_K;
const int qk_off = seff * TILED_TILE_K;
const int nb_stride = NB * nkr;
const int nb_off = seff * NB;
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
const int bs_off = seff * TILED_MICRO;
float acc[TILED_MICRO][TILED_MICRO];
memset(acc, 0, sizeof(acc));
// subdots at subblock granularity over the 256-K block, exact integer math
for (int s = 0; s < NB; s++) {
for (int j = 0; j < TILED_MICRO; j++) {
const int br = j0 + j;
const int8_t * q1 = &src1.q[br * qk_stride + qk_off + s * SUBBLK];
int32_t bsum = 0;
for (int u = 0; u < NS; u++) {
bsum += src1.bsums[(bs_off + s * NS + u) * bs_stride + br];
}
for (int i = 0; i < TILED_MICRO; i++) {
const int ar = i0 + i;
const int d_off = seff * TILED_MICRO + ar;
const uint8_t * q0 = &src0.q[ar * qk_stride + qk_off + s * SUBBLK];
int32_t raw = 0;
for (int e = 0; e < SUBBLK; e++) {
raw += (int32_t) q0[e] * (int32_t) q1[e];
}
// BIAS: subtract BIAS*bsum (src1's per-subblock code sum) from the exact int raw
int32_t corr = raw;
if constexpr (BIAS != 0) {
corr -= BIAS * bsum;
}
const int32_t scales_raw = (int32_t) src0.scales[ar * nb_stride + nb_off + s] * corr;
// d is NOT applied here: it is constant over the s-loop, so we apply it last before write-out
if constexpr (HAS_MIN) {
const int32_t mins_bsum = (int32_t) src0.mins[ar * nb_stride + nb_off + s] * bsum;
acc[i][j] += (float) src0.d[d_off] * (float) scales_raw
- (float) src0.dmin[d_off] * (float) mins_bsum;
} else {
acc[i][j] += (float) src0.d[d_off] * (float) scales_raw;
}
}
}
}
// Apply d and write out to buf
for (int i = 0; i < TILED_MICRO; i++) {
for (int j = 0; j < TILED_MICRO; j++) {
buf[(i0 + i) * buf_stride + (j0 + j)] += src1.d[seff * TILED_MICRO + j0 + j] * acc[i][j];
}
}
}
#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && defined(__AVX512DQ__)
// VNNI microkernel
// 8x16 band pass: 8 src0 rows (i0, i0+8) x 16 src1 cols (j0, j0+16)
// Math: s1_acc += scales_s*raw_s, s2_acc += mins_s*bsums_s per subblock; the row result
// is d0*s1_acc - dmin*s2_acc (f32, exact in range), applied once per row.
//
// Register pressure: the band pass holds acc16 + s1_acc = 16 zmm for the
// 8-row band (s2_acc has no dpbusd dependency, so it is computed after the
// band pass and adds no register pressure);
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
static void tiled_run_micro_vnni_8x16(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
constexpr int NB = TILED_TILE_K / SUBBLK;
constexpr int NS = SUBBLK / 16;
constexpr int NG = SUBBLK / 4;
// band width, see the register-pressure note above
constexpr int NUM_ROWS = 8;
// num_k = K-blocks per row (the narrow path holds num_k slabs at row stride num_k*256 so
// each weight row is one long stream); slab = this call's slab. num_k=1, slab=0 = standard.
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
// version serves all of it.
const int nkr = (NK > 0) ? NK : num_k;
const int seff = (NK > 0) ? 0 : slab;
const int qk_stride = nkr * TILED_TILE_K;
const int qk_off = seff * TILED_TILE_K;
const int nb_stride = NB * nkr;
const int nb_off = seff * NB;
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
const int bs_off = seff * TILED_MICRO;
const __m512 d1_vec = _mm512_load_ps(&src1.d[seff * TILED_MICRO + j0]);
__m512i s1_acc[NUM_ROWS];
for (int t = 0; t < NUM_ROWS; t++) { s1_acc[t] = _mm512_setzero_si512(); }
__m512i s2_acc[NUM_ROWS];
if constexpr (HAS_MIN) {
for (int t = 0; t < NUM_ROWS; t++) { s2_acc[t] = _mm512_setzero_si512(); }
}
const int32_t * bsums = src1.bsums; // int32 per-16 sums; NS > 1 combines NS lanes per subblock
for (int s = 0; s < NB; s++) {
__m512i bsums32 = _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS) * bs_stride + j0]);
for (int u = 1; u < NS; u++) {
bsums32 = _mm512_add_epi32(bsums32, _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS + u) * bs_stride + j0]));
}
__m512i bias32 = _mm512_setzero_si512();
if constexpr (BIAS != 0) {
bias32 = _mm512_mullo_epi32(bsums32, _mm512_set1_epi32(BIAS));
}
__m512i acc16[NUM_ROWS];
for (int t = 0; t < NUM_ROWS; t++) { acc16[t] = _mm512_setzero_si512(); }
#ifdef __GNUC__
#pragma GCC unroll 8 // pragma unrolled justified by measuing with/without
#endif
for (int g = 0; g < NG; g++) {
const int kg = s * NG + g;
// in-place interleave layout: [kg%16 @ TILED_TILE_K][kg/16 @ 64][row @ 4]
const __m512i codes = _mm512_load_si512((const __m512i *) &src1.q[(j0 / TILED_MICRO) * (TILED_MICRO * qk_stride)
+ qk_off + (kg % TILED_MICRO) * qk_stride + (kg / TILED_MICRO) * (TILED_MICRO * 4)]);
for (int t = 0; t < NUM_ROWS; t++) {
const uint32_t u4 = *(const uint32_t *) &src0.q[(i0 + t) * qk_stride + qk_off + kg * 4];
const __m512i u4b = _mm512_set1_epi32((int) u4);
acc16[t] = _mm512_dpbusd_epi32(acc16[t], u4b, codes);
}
}
// int correction: s1_acc += scales*(raw-BIAS*bsums)
for (int t = 0; t < NUM_ROWS; t++) {
const int ar = i0 + t;
__m512i rawi = acc16[t];
if constexpr (BIAS != 0) {
rawi = _mm512_sub_epi32(rawi, bias32);
}
s1_acc[t] = _mm512_add_epi32(s1_acc[t],
_mm512_mullo_epi32(rawi, _mm512_set1_epi32(src0.scales[ar * nb_stride + nb_off + s])));
}
}
// s2_acc += mins*bsums; independent of the dpbusd results, so it runs here instead of in the
// band pass: the band pass keeps the register budget for the 8-row band
if constexpr (HAS_MIN) {
for (int s = 0; s < NB; s++) {
__m512i bsums32 = _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS) * bs_stride + j0]);
for (int u = 1; u < NS; u++) {
bsums32 = _mm512_add_epi32(bsums32, _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS + u) * bs_stride + j0]));
}
for (int t = 0; t < NUM_ROWS; t++) {
s2_acc[t] = _mm512_add_epi32(s2_acc[t],
_mm512_mullo_epi32(bsums32, _mm512_set1_epi32(src0.mins[(i0 + t) * nb_stride + nb_off + s])));
}
}
}
// epilogue: int->float, apply per-row scales, store to buf
for (int t = 0; t < NUM_ROWS; t++) {
const int ar = i0 + t;
const int d_off = seff * TILED_MICRO + ar;
__m512 f1 = _mm512_cvtepi32_ps(s1_acc[t]);
__m512 result = _mm512_mul_ps(f1, _mm512_set1_ps(src0.d[d_off]));
if constexpr (HAS_MIN) {
__m512 f2 = _mm512_cvtepi32_ps(s2_acc[t]);
result = _mm512_fnmadd_ps(_mm512_set1_ps(src0.dmin[d_off]), f2, result);
}
float * p = &buf[(i0 + t) * buf_stride + j0];
_mm512_store_ps(p, _mm512_add_ps(_mm512_load_ps(p), _mm512_mul_ps(result, d1_vec)));
}
}
// 16x16 microtile as two explicit 8x16 band passes
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
static void tiled_run_microtile_vnni(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
tiled_run_micro_vnni_8x16<SUBBLK, HAS_MIN, BIAS, NK>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
tiled_run_micro_vnni_8x16<SUBBLK, HAS_MIN, BIAS, NK>(src0, src1, i0 + 8, j0, num_k, slab, buf, buf_stride);
}
#endif // __AVX512VNNI__ && __AVX512VL__
#if defined(__AVX2__)
// AVX2 kernel.
// We're effectively applying the existing vec_dot algorithms to an 8x16 block here.
// Different paths based on subblock size as it affects when/where we multiply in scales and apply mins
// SUBBLK=32: one 256-bit maddubs per (s, column), a 256-bit set1 scale,
// one 8-lane i32 accumulator per column.
// SUBBLK=16: two subblocks per 32B load. The 256-bit maddubs product
// still gives 16 i16 lanes (0..7 = subblock sp, 8..15 = sp+1), but the
// scale is applied per 128-bit half.
// With BIAS != 0 the q0 codes are biased (up to 241) and a maddubs i16 pair
// overflows, so the weight is debiased (w = q0 - BIAS, small). Two MACs:
// ACTBIAS: activation pre-biased (+128) by repack, one maddubs(q1b, w); the +128 is
// corrected out in the epilogue (128*sum_s scales[s]*w_bsum[s], precomputed in repack_src0).
// Fits only when the debiased weight magnitude Wmax <= 64 (2*255*Wmax < 32767).
// else (iq4_xs, Wmax = 127): the signed-int8 sign trick, ax = |w|, sy = q1*sign(w),
// one maddubs (2*127*127 < 32767) so it always fits.
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK, bool ACTBIAS>
static void tiled_run_microtile_avx2(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
constexpr int NB = TILED_TILE_K / SUBBLK;
constexpr int GROUP = 8; // Process 8 src1 columns concurrently (1 YMM vector)
static_assert(SUBBLK == 16 || SUBBLK == 32, "unsupported SUBBLK");
// num_k/slab: the narrow path holds num_k slabs at row stride num_k*256 (num_k=1, slab=0 = standard)
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
// version serves all of it.
const int nkr = (NK > 0) ? NK : num_k;
const int seff = (NK > 0) ? 0 : slab;
const int qk_stride = nkr * TILED_TILE_K;
const int qk_off = seff * TILED_TILE_K;
const int nb_stride = NB * nkr;
const int nb_off = seff * NB;
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
const int bs_off = seff * TILED_MICRO;
// Horizontal sum helper: converts 8x32-bit int lane inside a YMM register to scalar int32
auto hsum256_epi32 = [](const __m256i v) -> int32_t {
__m128i low = _mm256_castsi256_si128(v);
__m128i high = _mm256_extracti128_si256(v, 1);
__m128i sum128 = _mm_add_epi32(low, high);
sum128 = _mm_add_epi32(sum128, _mm_shuffle_epi32(sum128, _MM_SHUFFLE(2, 3, 0, 1)));
sum128 = _mm_add_epi32(sum128, _mm_shuffle_epi32(sum128, _MM_SHUFFLE(1, 0, 3, 2)));
return _mm_cvtsi128_si32(sum128);
};
for (int i = 0; i < TILED_MICRO; i++) {
const int ar = i0 + i;
const int d_off = seff * TILED_MICRO + ar;
const float d0 = src0.d[d_off];
const float dmin0 = src0.dmin[d_off];
const uint8_t * q0 = &src0.q[ar * qk_stride + qk_off];
const int32_t * scales_row = &src0.scales[ar * nb_stride + nb_off];
const int32_t * mins_row = &src0.mins[ar * nb_stride + nb_off];
for (int g = 0; g < TILED_MICRO; g += GROUP) {
// Pre-calculate weight pointers for this column group
const int8_t * q1_ptr[GROUP];
for (int t = 0; t < GROUP; t++) {
q1_ptr[t] = &src1.q[(j0 + g + t) * qk_stride + qk_off];
}
__m256i s2 = _mm256_setzero_si256();
__m256i acc[GROUP];
for (int t = 0; t < GROUP; t++) {
acc[t] = _mm256_setzero_si256();
}
if constexpr (SUBBLK == 32) {
for (int s = 0; s < NB; s++) {
// OPTIMIZATION 1: Use aligned loads (_mm256_load_si256)
const __m256i q0_32 = _mm256_load_si256((const __m256i *) &q0[s * SUBBLK]);
const __m256i scales16 = _mm256_set1_epi16(scales_row[s]);
if constexpr (BIAS != 0) {
const __m256i w = _mm256_sub_epi8(q0_32, _mm256_set1_epi8((int8_t) BIAS)); // debiased weight
if constexpr (ACTBIAS) {
// activation pre-biased (+128) by repack; q1b unsigned, w signed
#pragma GCC unroll 8
for (int t = 0; t < GROUP; t++) {
const __m256i q1b_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][s * SUBBLK]);
acc[t] = _mm256_add_epi32(acc[t],
_mm256_madd_epi16(scales16, _mm256_maddubs_epi16(q1b_32, w)));
}
} else {
// iq4_xs (Wmax > 64 overflows act-bias): sign trick, both operands <= 127
const __m256i ax = _mm256_sign_epi8(w, w);
#pragma GCC unroll 8
for (int t = 0; t < GROUP; t++) {
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][s * SUBBLK]);
acc[t] = _mm256_add_epi32(acc[t],
_mm256_madd_epi16(scales16, _mm256_maddubs_epi16(ax, _mm256_sign_epi8(q1_32, w))));
}
}
} else {
#pragma GCC unroll 8
for (int t = 0; t < GROUP; t++) {
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][s * SUBBLK]);
acc[t] = _mm256_add_epi32(acc[t],
_mm256_madd_epi16(scales16, _mm256_maddubs_epi16(q0_32, q1_32)));
}
}
if constexpr (HAS_MIN) {
const __m256i bsums_v = _mm256_add_epi32(
_mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + s * 2) * bs_stride + j0 + g]),
_mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + s * 2 + 1) * bs_stride + j0 + g]));
s2 = _mm256_add_epi32(s2, _mm256_mullo_epi32(bsums_v, _mm256_set1_epi32(mins_row[s])));
}
}
} else { // SUBBLK == 16
for (int sp = 0; sp < NB; sp += 2) {
const __m256i q0_32 = _mm256_load_si256((const __m256i *) &q0[sp * SUBBLK]);
const __m256i scalesv = _mm256_set_m128i(_mm_set1_epi16(scales_row[sp + 1]), _mm_set1_epi16(scales_row[sp]));
if constexpr (BIAS != 0) {
const __m256i w = _mm256_sub_epi8(q0_32, _mm256_set1_epi8((int8_t) BIAS)); // debiased weight
if constexpr (ACTBIAS) {
// activation pre-biased (+128) by repack; q1b unsigned, w signed
#pragma GCC unroll 8
for (int t = 0; t < GROUP; t++) {
const __m256i q1b_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][sp * SUBBLK]);
acc[t] = _mm256_add_epi32(acc[t], _mm256_madd_epi16(
scalesv, _mm256_maddubs_epi16(q1b_32, w)));
}
} else {
// iq4_xs (Wmax > 64 overflows act-bias): sign trick, both operands <= 127
const __m256i ax = _mm256_sign_epi8(w, w);
#pragma GCC unroll 8
for (int t = 0; t < GROUP; t++) {
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][sp * SUBBLK]);
acc[t] = _mm256_add_epi32(acc[t], _mm256_madd_epi16(
scalesv, _mm256_maddubs_epi16(ax, _mm256_sign_epi8(q1_32, w))));
}
}
} else {
#pragma GCC unroll 8
for (int t = 0; t < GROUP; t++) {
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][sp * SUBBLK]);
acc[t] = _mm256_add_epi32(acc[t], _mm256_madd_epi16(
scalesv, _mm256_maddubs_epi16(q0_32, q1_32)));
}
}
if constexpr (HAS_MIN) {
const __m256i bsums0_v = _mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + sp) * bs_stride + j0 + g]);
const __m256i bsums1_v = _mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + sp + 1) * bs_stride + j0 + g]);
s2 = _mm256_add_epi32(s2, _mm256_add_epi32(
_mm256_mullo_epi32(bsums0_v, _mm256_set1_epi32(mins_row[sp])),
_mm256_mullo_epi32(bsums1_v, _mm256_set1_epi32(mins_row[sp + 1]))));
}
}
}
// OPTIMIZATION 2: Vectorized Epilogue with FMA3 gating
// Reduce accumulators into a 256-bit vector across 8 columns
alignas(32) int32_t s1_vals[GROUP];
#pragma GCC unroll 8
for (int t = 0; t < GROUP; t++) {
s1_vals[t] = hsum256_epi32(acc[t]);
}
__m256i s1_vec = _mm256_load_si256((const __m256i *) s1_vals);
if constexpr (BIAS != 0 && ACTBIAS) {
// act-bias correction: raw = desired + 128*sum_s scales[s]*w_bsum[s]; corr precomputed in repack_src0
s1_vec = _mm256_sub_epi32(s1_vec, _mm256_set1_epi32(mins_row[0]));
}
const __m256 src1_d_vec = _mm256_load_ps(&src1.d[seff * TILED_MICRO + j0 + g]);
__m256 res_vec = _mm256_mul_ps(_mm256_set1_ps(d0), _mm256_cvtepi32_ps(s1_vec));
if constexpr (HAS_MIN) {
#if defined(__FMA__)
// FMA3 implementation: res_vec = res_vec - (dmin0 * s2)
res_vec = _mm256_fnmadd_ps(_mm256_set1_ps(dmin0), _mm256_cvtepi32_ps(s2), res_vec);
#else
// Non-FMA fallback
res_vec = _mm256_sub_ps(res_vec, _mm256_mul_ps(_mm256_set1_ps(dmin0), _mm256_cvtepi32_ps(s2)));
#endif
}
float * buf_ptr = &buf[ar * buf_stride + j0 + g];
__m256 current_buf = _mm256_loadu_ps(buf_ptr);
#if defined(__FMA__)
// FMA3 accumulation: current_buf + (res_vec * src1_d_vec)
__m256 updated_buf = _mm256_fmadd_ps(res_vec, src1_d_vec, current_buf);
#else
__m256 updated_buf = _mm256_add_ps(current_buf, _mm256_mul_ps(res_vec, src1_d_vec));
#endif
_mm256_storeu_ps(buf_ptr, updated_buf);
}
}
}
#endif // __AVX2__
#if defined(__AVX__) && !defined(__AVX2__)
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
static void tiled_run_microtile_avx(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
constexpr int NB = TILED_TILE_K / SUBBLK;
constexpr int NS = SUBBLK / 16;
constexpr int GROUP = 4; // src1 columns per group: one acc32 per column
// num_k/slab: the narrow path holds num_k slabs at row stride num_k*256 (num_k=1, slab=0 = standard)
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
// version serves all of it.
const int nkr = (NK > 0) ? NK : num_k;
const int seff = (NK > 0) ? 0 : slab;
const int qk_stride = nkr * TILED_TILE_K;
const int qk_off = seff * TILED_TILE_K;
const int nb_stride = NB * nkr;
const int nb_off = seff * NB;
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
const int bs_off = seff * TILED_MICRO;
for (int i = 0; i < TILED_MICRO; i++) {
const int ar = i0 + i;
const int d_off = seff * TILED_MICRO + ar;
const float d0 = src0.d[d_off];
const float dmin0 = src0.dmin[d_off];
const uint8_t * q0 = &src0.q[ar * qk_stride + qk_off];
const int32_t * scales_row = &src0.scales[ar * nb_stride + nb_off];
const int32_t * mins_row = &src0.mins[ar * nb_stride + nb_off];
for (int g = 0; g < TILED_MICRO; g += GROUP) {
const int8_t * q1g[GROUP];
__m128i acc[GROUP];
__m128i s2 = _mm_setzero_si128(); // sum_s mins_s * bsums_s
for (int t = 0; t < GROUP; t++) {
q1g[t] = &src1.q[(j0 + g + t) * qk_stride + qk_off];
acc[t] = _mm_setzero_si128();
}
for (int s = 0; s < NB; s++) {
const __m128i scales16 = _mm_set1_epi16(scales_row[s]);
if constexpr (HAS_MIN) {
__m128i bsums_v = _mm_setzero_si128(); // per-col bsums as a 4-lane i32 vector
for (int u = 0; u < NS; u++) {
bsums_v = _mm_add_epi32(bsums_v, _mm_loadu_si128(
(const __m128i *) &src1.bsums[(bs_off + s * NS + u) * bs_stride + j0 + g]));
}
s2 = _mm_add_epi32(s2, _mm_mullo_epi32(bsums_v, _mm_set1_epi32(mins_row[s])));
}
for (int u = 0; u < NS; u++) {
const __m128i a16 = _mm_loadu_si128((const __m128i *) &q0[s * SUBBLK + u * 16]);
if constexpr (BIAS != 0) {
// debiased weight w = a16 - BIAS; signed dot via the sign trick (see header)
const __m128i a16s = _mm_sub_epi8(a16, _mm_set1_epi8((int8_t) BIAS));
const __m128i ax = _mm_sign_epi8(a16s, a16s);
for (int t = 0; t < GROUP; t++) {
const __m128i q1_16 = _mm_loadu_si128((const __m128i *) &q1g[t][s * SUBBLK + u * 16]);
acc[t] = _mm_add_epi32(acc[t],
_mm_madd_epi16(scales16, _mm_maddubs_epi16(ax, _mm_sign_epi8(q1_16, a16s))));
}
} else {
for (int t = 0; t < GROUP; t++) {
acc[t] = _mm_add_epi32(acc[t], _mm_madd_epi16(scales16,
_mm_maddubs_epi16(a16,
_mm_loadu_si128((const __m128i *) &q1g[t][s * SUBBLK + u * 16]))));
}
}
}
}
int32_t s2_s[GROUP] = { 0, 0, 0, 0 };
if constexpr (HAS_MIN) {
_mm_storeu_si128((__m128i *) s2_s, s2);
}
for (int t = 0; t < GROUP; t++) {
// 4 i32 lanes -> scalar
__m128i v = _mm_shuffle_epi32(acc[t], _MM_SHUFFLE(2, 3, 0, 1));
acc[t] = _mm_add_epi32(acc[t], v);
v = _mm_shuffle_epi32(acc[t], _MM_SHUFFLE(1, 0, 3, 2));
acc[t] = _mm_add_epi32(acc[t], v);
const int32_t s1 = _mm_cvtsi128_si32(acc[t]);
float res = d0 * (float) s1;
if constexpr (HAS_MIN) {
res -= dmin0 * (float) s2_s[t];
}
buf[ar * buf_stride + j0 + g + t] += src1.d[seff * TILED_MICRO + j0 + g + t] * res;
}
}
}
}
#endif // __AVX__ && !__AVX2__
// main microtile entry point. num_k = K-blocks per row (the narrow path reads num_k slabs at row
// stride num_k*256 so each weight row is one long stream); slab = this call's slab. num_k=1,
// slab=0 is the standard single-slab path.
// num_k==1 dispatches to the NK=1 instantiation, whose strides/offsets are compile-time
// constants (shifts/scaled-leas). num_k>1 (the narrow path) uses the NK=0 runtime-strided
// version; it is memory-bound, so one version serves all of it.
template <int SUBBLK, bool HAS_MIN, int BIAS, bool ACTBIAS>
void tiled_run_microtile(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
const bool standard = (num_k == 1);
#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && defined(__AVX512DQ__)
if (standard) {
tiled_run_microtile_vnni<SUBBLK, HAS_MIN, BIAS, 1>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
} else {
tiled_run_microtile_vnni<SUBBLK, HAS_MIN, BIAS, 0>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
}
#elif defined(__AVX2__)
if (standard) {
tiled_run_microtile_avx2<SUBBLK, HAS_MIN, BIAS, 1, ACTBIAS>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
} else {
tiled_run_microtile_avx2<SUBBLK, HAS_MIN, BIAS, 0, ACTBIAS>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
}
#elif defined(__AVX__)
if (standard) {
tiled_run_microtile_avx<SUBBLK, HAS_MIN, BIAS, 1>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
} else {
tiled_run_microtile_avx<SUBBLK, HAS_MIN, BIAS, 0>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
}
#else
if (standard) {
tiled_run_microtile_scalar<SUBBLK, HAS_MIN, BIAS, 1>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
} else {
tiled_run_microtile_scalar<SUBBLK, HAS_MIN, BIAS, 0>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
}
#endif
}
// explicit instantiations for the in-use formats (q4_K and q5_K share the constants)
// ACTBIAS selects the act-bias MAC (Wmax <= 64); iq4_xs (Wmax 127) keeps the sign trick
// q4_K / q5_K: BIAS = 0
template void tiled_run_microtile<32, true, 0, false>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
// iq4_xs: sign trick (Wmax 127 overflows act-bias)
template void tiled_run_microtile<32, false, 128, false>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
// iq2_xxs, iq3_xxs, iq3_s, iq1_s: act-bias (Wmax <= 62)
template void tiled_run_microtile<32, false, 128, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
// iq2_xs, iq2_s, iq1_m: act-bias
template void tiled_run_microtile<16, false, 128, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
// q6_K: act-bias (Wmax 32)
template void tiled_run_microtile<16, false, 32, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
// q3_K: act-bias (Wmax 4)
template void tiled_run_microtile<16, false, 4, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
// q2_K: BIAS = 0
template void tiled_run_microtile<16, true, 0, false>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
#define MIN(a, b) ((a) < (b) ? (a) : (b))
// block_q8_K in 4-byte words, for the int32 gather indices
static_assert(sizeof(block_q8_K) == 292 && offsetof(block_q8_K, qs) == 4,
"block_q8_K layout changed, fix the src1 repack");
// VNNI: interleave the natural [row][k] codes in-place into the group-local
// [kg%16][kg/16][row][4] layout. base points at row 0 (row r at base + r*row_stride);
// n_tiles 16x16 int32 tiles are transposed to the right (tile t occupies bytes
// [t*64, t*64+64) per row). Each 16-row group is self-contained in row_stride*16 bytes.
#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && defined(__AVX512DQ__)
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias) {
GGML_UNUSED(bias);
const int n_tiles = num_k * (TILED_TILE_K / (TILED_MICRO * 4));
const int row_stride = num_k * TILED_TILE_K;
uint8_t * base = (uint8_t *) src1->q + row0 * row_stride;
for (int t = 0; t < n_tiles; t++) {
uint8_t * cbase = base + t * (TILED_MICRO * 4);
__m512i v[16];
for (int r = 0; r < 16; r++) {
v[r] = _mm512_load_si512((const __m512i *) (cbase + r * row_stride));
}
// 16x16 int32 transpose, 4 butterfly phases
// Phase 1: 1-element interleave, pairs (0,1), (2,3), ..., (14,15)
{
const __m512i idx_a = _mm512_setr_epi32(0, 16, 1, 17, 2, 18, 3, 19, 4, 20, 5, 21, 6, 22, 7, 23);
const __m512i idx_b = _mm512_setr_epi32(8, 24, 9, 25, 10, 26, 11, 27, 12, 28, 13, 29, 14, 30, 15, 31);
for (int i = 0; i < 16; i += 2) {
__m512i a = v[i], b = v[i+1];
v[i] = _mm512_permutex2var_epi32(a, idx_a, b);
v[i+1] = _mm512_permutex2var_epi32(a, idx_b, b);
}
}
// Phase 2: 2-element interleave, pairs (0,2), (1,3), (4,6), (5,7), ...
{
const __m512i idx_a = _mm512_setr_epi32(0, 1, 16, 17, 4, 5, 20, 21, 8, 9, 24, 25, 12, 13, 28, 29);
const __m512i idx_b = _mm512_setr_epi32(2, 3, 18, 19, 6, 7, 22, 23, 10, 11, 26, 27, 14, 15, 30, 31);
for (int i = 0; i < 16; i += 4) {
__m512i a = v[i], b = v[i+2];
v[i] = _mm512_permutex2var_epi32(a, idx_a, b);
v[i+2] = _mm512_permutex2var_epi32(a, idx_b, b);
a = v[i+1], b = v[i+3];
v[i+1] = _mm512_permutex2var_epi32(a, idx_a, b);
v[i+3] = _mm512_permutex2var_epi32(a, idx_b, b);
}
}
// Phase 3: 4-element interleave, pairs (0,4), (1,5), ..., (7,11), (8,12), ...
{
const __m512i idx_a = _mm512_setr_epi32(0, 1, 2, 3, 16, 17, 18, 19, 4, 5, 6, 7, 20, 21, 22, 23);
const __m512i idx_b = _mm512_setr_epi32(8, 9, 10, 11, 24, 25, 26, 27, 12, 13, 14, 15, 28, 29, 30, 31);
for (int i = 0; i < 16; i += 8) {
for (int j = 0; j < 4; j++) {
__m512i a = v[i+j], b = v[i+4+j];
v[i+j] = _mm512_permutex2var_epi32(a, idx_a, b);
v[i+4+j] = _mm512_permutex2var_epi32(a, idx_b, b);
}
}
}
// Phase 4: 8-element interleave, pairs (0,8), (1,9), ..., (7,15)
{
const __m512i idx_a = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 16, 17, 18, 19, 20, 21, 22, 23);
const __m512i idx_b = _mm512_setr_epi32(8, 9, 10, 11, 12, 13, 14, 15, 24, 25, 26, 27, 28, 29, 30, 31);
for (int i = 0; i < 8; i++) {
__m512i a = v[i], b = v[i+8];
v[i] = _mm512_permutex2var_epi32(a, idx_a, b);
v[i+8] = _mm512_permutex2var_epi32(a, idx_b, b);
}
}
// store: v[g] holds 16 int32s for k-group col_order[g], rows 0..15
static const int col_order[16] = {0, 8, 1, 9, 4, 12, 5, 13, 2, 10, 3, 11, 6, 14, 7, 15};
for (int g = 0; g < 16; g++) {
_mm512_store_si512((void *) (cbase + col_order[g] * row_stride), v[g]);
}
}
}
#elif defined(__AVX2__)
// act-bias: pre-bias the activation in place (+128) so the MAC's maddubs(q1b, w) sees unsigned bytes;
// no transpose on AVX2 (natural layout). 16 rows x n_tiles*64 bytes, row r at base + r*row_stride
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias) {
if (!bias) {
return;
}
const int row_stride = num_k * TILED_TILE_K;
uint8_t * base = (uint8_t *) src1->q + row0 * row_stride;
const int nbytes = row_stride; // full row width
const __m256i b128 = _mm256_set1_epi8((int8_t) 0x80);
for (int r = 0; r < TILED_MICRO; r++) {
uint8_t * row = base + r * row_stride;
for (int off = 0; off < nbytes; off += 32) {
const __m256i v = _mm256_loadu_si256((const __m256i *) (row + off));
_mm256_storeu_si256((__m256i *) (row + off), _mm256_xor_si256(v, b128));
}
}
}
#else
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias) {
GGML_UNUSED(src1); GGML_UNUSED(row0); GGML_UNUSED(num_k); GGML_UNUSED(bias);
}
#endif
#if defined(__AVX2__)
// sum 32 unsigned bytes to int32 (via maddubs -> int16 pairs -> hsum)
static int32_t tiled_byte_sum_32(const uint8_t * p) {
const __m256i dot = _mm256_maddubs_epi16(_mm256_loadu_si256((const __m256i *) p), _mm256_set1_epi8(1));
const __m256i sum = _mm256_madd_epi16(_mm256_set1_epi16(1), dot); // 8 int32
// reduce 8 lanes: lo+hi -> 4 lanes, hadd -> 2, hadd -> 1 (a single hadd pair only folds 4 lanes)
const __m128i t = _mm_add_epi32(_mm256_castsi256_si128(sum), _mm256_extracti128_si256(sum, 1));
const __m128i h = _mm_hadd_epi32(t, t);
return _mm_cvtsi128_si32(_mm_hadd_epi32(h, h));
}
// sum 16 unsigned bytes to int32
static int32_t tiled_byte_sum_16(const uint8_t * p) {
const __m128i dot = _mm_maddubs_epi16(_mm_loadu_si128((const __m128i *) p), _mm_set1_epi8(1));
const __m128i sum = _mm_madd_epi16(_mm_set1_epi16(1), dot); // 4 int32
const __m128i h = _mm_hadd_epi32(sum, sum);
return _mm_cvtsi128_si32(_mm_hadd_epi32(h, h));
}
#endif
// act-bias: precompute the per-weight-row correction corr = 128*sum_s scales[s]*w_bsum[s] where
// w_bsum[s] = sum of the debiased weight bytes in subblock s. Stored in mins[r][0] (unused for
// HAS_MIN = false, which is exactly the act-bias case). Reused across all act bands and weight groups.
// No-op off AVX2 (the act-bias MAC is AVX2-only) and when corr is false (iq4_xs / BIAS = 0).
template <int SUBBLK>
void tiled_repack_src0(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr) {
#if defined(__AVX2__)
if (!corr) {
return;
}
constexpr int NB = TILED_TILE_K / SUBBLK;
const int nb_stride = NB * num_k;
for (int slab = 0; slab < num_k; slab++) {
for (int r = 0; r < n_rows; r++) {
const uint8_t * q = &tile->q[r * (num_k * TILED_TILE_K) + slab * TILED_TILE_K];
const int32_t * scales = &tile->scales[r * nb_stride + slab * NB];
int32_t corr_val = 0;
for (int s = 0; s < NB; s++) {
const int32_t qsum = (SUBBLK == 32) ? tiled_byte_sum_32(&q[s * SUBBLK]) : tiled_byte_sum_16(&q[s * SUBBLK]);
corr_val += 128 * scales[s] * (qsum - SUBBLK * BIAS);
}
tile->mins[r * nb_stride + slab * NB] = corr_val;
}
}
#else
GGML_UNUSED(tile); GGML_UNUSED(n_rows); GGML_UNUSED(num_k); GGML_UNUSED(BIAS); GGML_UNUSED(corr);
#endif
}
template void tiled_repack_src0<16>(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr);
template void tiled_repack_src0<32>(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr);
+180
View File
@@ -0,0 +1,180 @@
#pragma once
// Tiled matmul kernel API: tile structs, kernel definitions
// Currently only optimized for x86, new architectures should implement:
// tiled_run_microtile: 16x16 microkernel
// tiled_repack_src0: Optional repack/recalculation of src0, per macrotile
// tiled_repack_src1: Optional repack/recalculation of src1, per microtile-band
// bit unpacking routines: tiled_unpk_nib4, tiled_unpk_2bit, tiled_unpk_or
// LUT value expansion routines: tiled_lut8, tiled_unpk_sign32, tiled_unpk_tern8
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include <stddef.h>
#include <stdint.h>
#if defined(__AVX2__)
#include <immintrin.h>
#endif
#define TILED_TILE_K 256 // one QK_K block
#define TILED_TILE_ROWS 256 // max window rows, ragged at edges
#define TILED_MICRO 16 // microtile edge (also the bsums code-sum granularity)
#define TILED_WS_SLOT (512 * 1024) // per-thread workspace slot, a clean 512KB (multiple of 64B)
// src0 tile: weight side, shared by all formats.
// scales/mins are sized for the max subblock count (SUBBLK=16);
// SUBBLK=32 formats index at stride 8 and leave the slack unused.
struct tiled_tile_src0 {
static constexpr int NB_MAX = TILED_TILE_K / 16; // max subblocks per 256-elem block
alignas(64) uint8_t q[TILED_TILE_ROWS * TILED_TILE_K]; // unsigned quants, widened to uint8
float d[TILED_TILE_ROWS]; // One d from each input block, widened to f32
float dmin[TILED_TILE_ROWS]; // dmin from each input block (if applicable), widened to F32
int32_t scales[TILED_TILE_ROWS * NB_MAX]; // per-subblock scale, stored as int32_t
int32_t mins[TILED_TILE_ROWS * NB_MAX]; // per-subblock min, used when HAS_MIN
};
// src1 tile: built from q8_K (wdata)
struct tiled_tile_src1 {
// q8 codes, one byte per element. Note for VNNI these are reshaped + transposed to be suitable for dpbusd.
alignas(64) int8_t q[TILED_TILE_ROWS * TILED_TILE_K];
// per-16 code sums from q8_k (int16), widened to int32 so the kernels load them directly, no per-use cvt
alignas(64) int32_t bsums[(TILED_TILE_K / 16) * TILED_TILE_ROWS];
// f32 (not f16): q8_k stores fp16, the unpack converts once
float d[TILED_TILE_ROWS];
};
// per-thread workspace: all tiled state lives here, allocated in wdata (one slot per thread)
struct tiled_ws {
tiled_tile_src0 src0;
tiled_tile_src1 src1;
alignas(64) float acc[TILED_TILE_ROWS * TILED_TILE_ROWS];
};
static_assert(sizeof(tiled_ws) <= TILED_WS_SLOT, "tiled workspace exceeds the 512KB per-thread slot");
// the slot base is 64B-aligned and TILED_WS_SLOT is a multiple of 64B, so each per-thread slot
// is 64B-aligned; these pin the tile fields and the acc buffer at aligned offsets within a slot
static_assert(offsetof(tiled_ws, src0) % 64 == 0, "src0 not 64B-aligned in the workspace");
static_assert(offsetof(tiled_ws, src1) % 64 == 0, "src1 not 64B-aligned in the workspace");
static_assert(offsetof(tiled_ws, acc) % 64 == 0, "acc not 64B-aligned in the workspace");
// unpack primitives for reading quants, defined as inline here to keep arch-specific code in kernel.h/.cpp
// If this section gets too hairy later, we can break up into separate includes.
#if defined(__AVX2__)
// packed 4-bit codes -> low nibbles (lo) + high nibbles (hi)
inline void tiled_unpk_nib4(const uint8_t * src, uint8_t * lo, uint8_t * hi) {
const __m256i v = _mm256_loadu_si256((const __m256i *) src);
// mask before the lane shift so bits do not cross byte boundaries
_mm256_storeu_si256((__m256i *) lo, _mm256_and_si256(v, _mm256_set1_epi8(0x0F)));
_mm256_storeu_si256((__m256i *) hi, _mm256_srli_epi32(_mm256_and_si256(v, _mm256_set1_epi8((int8_t) 0xF0)), 4));
}
// 2-bit values at bit offset S
template <int S> inline void tiled_unpk_2bit(const uint8_t * src, uint8_t * dst) {
_mm256_storeu_si256((__m256i *) dst, _mm256_and_si256(
_mm256_srli_epi32(_mm256_loadu_si256((const __m256i *) src), S), _mm256_set1_epi8(0x03)));
}
// OR the M-bit value at bit offset S of src into bit offset D of dst
template <int S, int D, int M>
inline void tiled_unpk_or(uint8_t * dst, const uint8_t * src) {
const __m256i v = _mm256_slli_epi32(_mm256_and_si256(
_mm256_srli_epi32(_mm256_loadu_si256((const __m256i *) src), S), _mm256_set1_epi8((uint8_t) M)), D);
_mm256_storeu_si256((__m256i *) dst, _mm256_or_si256(_mm256_loadu_si256((const __m256i *) dst), v));
}
// Unpacking kernels for IQ quants
// LUT value expansion for the LUT-based formats (iq4_xs, iq grids): the bit unpackers
// above give the indices, these expand 8/16 of them to widened codes in one pass
// 16-entry byte LUT: dst[j] = lut[src[j]] (16 bytes)
inline void tiled_lut8(const uint8_t * lut, const uint8_t * src, uint8_t * dst) {
_mm_storeu_si128((__m128i *) dst, _mm_shuffle_epi8(_mm_loadu_si128((const __m128i *) lut),
_mm_loadu_si128((const __m128i *) src)));
}
// 32 grid magnitudes (4 x 64-bit groups g0..g3, 8 values each) + 4 sign bytes
// (byte l signs values 8*l .. 8*l+7) -> codes stored as (value + 128)
inline void tiled_unpk_sign32(uint64_t g0, uint64_t g1, uint64_t g2, uint64_t g3,
const uint8_t signs[4], uint8_t * dst32) {
const __m256i v = _mm256_set_epi64x((int64_t) g3, (int64_t) g2, (int64_t) g1, (int64_t) g0);
const __m256i sv = _mm256_shuffle_epi8(_mm256_set1_epi32((int32_t) (signs[0] | signs[1] << 8 | signs[2] << 16 | signs[3] << 24)),
_mm256_setr_epi8(0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1,
2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3));
const __m256i sel = _mm256_set1_epi64x((int64_t) 0x8040201008040201ULL);
// 0xFF in every lane whose sign bit is set; GFNI does the and+compare in one instruction
#ifdef __GFNI__
const __m256i mask = _mm256_gf2p8affine_epi64_epi8(sel, sv, 0);
#else
const __m256i mask = _mm256_cmpeq_epi8(_mm256_and_si256(sv, sel), sel);
#endif
// v ^ mask - mask negates the signed lanes (v elsewhere); + 128 gives the biased code
const __m256i sgn = _mm256_sub_epi8(_mm256_xor_si256(v, mask), mask);
_mm256_storeu_si256((__m256i *) dst32, _mm256_add_epi8(sgn, _mm256_set1_epi8((int8_t) 128)));
}
// 8 ternary grid bytes (0 = 0, 1 = +1, 0xFF = -1): dst[j] = 128 + delta + 8 * (int8_t) src[j]
inline void tiled_unpk_tern8(const uint8_t * src, int8_t delta, uint8_t * dst) {
const __m128i v = _mm_cvtepi8_epi16(_mm_loadl_epi64((const __m128i *) src));
const __m128i p = _mm_add_epi16(_mm_slli_epi16(v, 3), _mm_set1_epi16(128 + (int) delta));
_mm_storel_epi64((__m128i *) dst, _mm_packus_epi16(p, _mm_setzero_si128()));
}
#else
// Scalar definitions for unpackers.
inline void tiled_unpk_nib4(const uint8_t * src, uint8_t * lo, uint8_t * hi) {
for (int l = 0; l < 32; l++) { lo[l] = (uint8_t) (src[l] & 0xF); hi[l] = (uint8_t) (src[l] >> 4); }
}
template <int S>
inline void tiled_unpk_2bit(const uint8_t * src, uint8_t * dst) {
for (int l = 0; l < 32; l++) { dst[l] = (uint8_t) ((src[l] >> S) & 3); }
}
template <int S, int D, int M>
inline void tiled_unpk_or(uint8_t * dst, const uint8_t * src) {
for (int l = 0; l < 32; l++) { dst[l] = (uint8_t) (dst[l] | (((src[l] >> S) & M) << D)); }
}
inline void tiled_lut8(const uint8_t * lut, const uint8_t * src, uint8_t * dst) {
for (int j = 0; j < 16; j++) { dst[j] = lut[src[j]]; }
}
inline void tiled_unpk_sign32(uint64_t g0, uint64_t g1, uint64_t g2, uint64_t g3,
const uint8_t signs[4], uint8_t * dst32) {
const uint64_t g[4] = { g0, g1, g2, g3 };
for (int l = 0; l < 4; l++) {
const uint8_t * v = (const uint8_t *) &g[l];
const uint8_t s = signs[l];
for (int j = 0; j < 8; j++) {
dst32[8 * l + j] = (s & (1 << j)) ? (uint8_t) (128 - v[j]) : (uint8_t) (128 + v[j]);
}
}
}
// 8 ternary grid bytes (0 = 0, 1 = +1, 0xFF = -1): dst[j] = 128 + delta + 8 * (int8_t) src[j]
inline void tiled_unpk_tern8(const uint8_t * src, int8_t delta, uint8_t * dst) {
for (int j = 0; j < 8; j++) { dst[j] = (uint8_t) (128 + (int) delta + 8 * (int8_t) src[j]); }
}
#endif
// Accumulate one 16x16 microtile (src0 rows [i0, i0+16), src1 cols [j0, j0+16))
// over one 256-K slab held in the tiles into a j-major float buffer
// (row width buf_stride): buf[i*buf_stride + j] += partial.
// SUBBLK/HAS_MIN/BIAS are the src0 format constants (see tiled_tile_src0).
// ACTBIAS (AVX2 only): the activation is pre-biased +128 by tiled_repack_src1,
// num_k = K-blocks per row: the tile holds num_k slabs at row stride num_k*256
// (Default case is num_k=1, 256x256 tiles, we go to longer num_k to improve memory bandwidth when num_rows is small)
template <int SUBBLK, bool HAS_MIN, int BIAS, bool ACTBIAS>
void tiled_run_microtile(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
// Optional repack, if profitable for the kernel.
// Repacks one 16-row band of src1 codes, called by driver as we reach each 16-row band in outer loop
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias);
// Optional repack, if profitable for the kernel
// Repack the entire src0 panel in place, called by the driver immediately after dequant
template <int SUBBLK>
void tiled_repack_src0(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr);
File diff suppressed because it is too large Load Diff
+31
View File
@@ -0,0 +1,31 @@
#pragma once
#include <stddef.h>
#include <stdint.h>
#ifdef __cplusplus
extern "C" {
#endif
// Amount of wdata to reserve for tiled workspaces
size_t ggml_tiled_wdata_size(int n_tasks, struct ggml_tensor * dst);
// tiled K-quant matmul; returns true if the op was computed here,
// false to fall through to the stock path
bool ggml_compute_forward_mul_mat_tiled(const struct ggml_compute_params * params,
struct ggml_tensor * dst);
// MUL_MAT_ID (MoE) path, one expert; returns true if the expert was computed here, per expert
// eligibility (type gate, batch floor) is decided inside. expert_rows points at the expert's
// row of the matrix_rows table of (expert slot, batch row) int32 pairs; scratch is the
// per-thread tiled_ws region reserved in wdata, n_tasks of ggml_tiled_ws_size() bytes
bool ggml_compute_forward_mul_mat_id_tiled(const struct ggml_compute_params * params,
struct ggml_tensor * dst,
int64_t cur_a,
int64_t cne1,
const int32_t * expert_rows,
char * scratch);
#ifdef __cplusplus
}
#endif
+13
View File
@@ -1,5 +1,18 @@
cmake_minimum_required(VERSION 3.18) # for CMAKE_CUDA_ARCHITECTURES
# ARM64EC uses the x64 ABI, so it needs the x64 import libraries. The arm64 ones
# hold native ARM64 symbols and cannot satisfy EC references at link time.
# FindCUDAToolkit picks the library dir from the host arch, not the target, so
# anchor its search at lib/x64 here. It derives everything else from CUDA_CUDART.
string(TOLOWER "${CMAKE_GENERATOR_PLATFORM}" GGML_CUDA_PLATFORM_LWR)
if (MSVC AND GGML_CUDA_PLATFORM_LWR STREQUAL "arm64ec")
find_library(CUDA_CUDART
NAMES cudart
HINTS ${CUDAToolkit_ROOT} ENV CUDA_PATH
PATH_SUFFIXES lib/x64
)
endif()
find_package(CUDAToolkit)
if (CUDAToolkit_FOUND)
+6 -15
View File
@@ -111,9 +111,9 @@
#define GGML_CUDA_CC_IS_QY2(cc) (cc >= GGML_CUDA_CC_QY2 && cc < GGML_CUDA_CC_PH1)
#define GGML_CUDA_CC_IS_PH1(cc) (cc >= GGML_CUDA_CC_PH1)
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
#if !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
# define GGML_CUDA_USE_CUB
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
#endif // !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
// PDL host-side support (cudaLaunchKernelEx) requires CUDART >= 11.8.
// However, this has been bugged in CTK < 12.3 for MSVC builds, see
@@ -237,7 +237,7 @@ static const char * cu_get_error_str(CUresult err) {
#define CU_CHECK(err) CUDA_CHECK_GEN(err, CUDA_SUCCESS, cu_get_error_str)
#endif
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
#if !defined(GGML_USE_HIP)
# define CUDA_SET_SHARED_MEMORY_LIMIT(kernel, nbytes) \
do { \
static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = { false }; \
@@ -252,7 +252,7 @@ static const char * cu_get_error_str(CUresult err) {
do { \
GGML_UNUSED(nbytes); \
} while (0)
#endif // !(defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
#endif // !defined(GGML_USE_HIP)
#if CUDART_VERSION >= 11010 || defined(GGML_USE_MUSA)
#define GGML_CUDA_ASSUME(x) __builtin_assume(x)
@@ -397,7 +397,7 @@ static constexpr __device__ int ggml_cuda_get_physical_warp_size() {
// Maximum number of bytes that can be copied in a single instruction.
static constexpr __device__ int ggml_cuda_get_max_cpy_bytes() {
#ifdef GGML_USE_HIP
#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
return 16;
#else
#if __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA
@@ -405,7 +405,7 @@ static constexpr __device__ int ggml_cuda_get_max_cpy_bytes() {
#else
return 8;
#endif // __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA
#endif // GGML_USE_HIP
#endif // defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
}
@@ -424,10 +424,6 @@ static __device__ void no_device_code(
__trap();
GGML_UNUSED(no_device_code); // suppress unused function warning
#if defined(GGML_USE_MUSA)
__builtin_unreachable();
#endif // defined(GGML_USE_MUSA)
}
#ifdef __CUDA_ARCH__
@@ -696,16 +692,11 @@ static __device__ __forceinline__ half2 ggml_cuda_hmax2(const half2 a, const hal
template<int width = WARP_SIZE>
static __device__ __forceinline__ half2 warp_reduce_max(half2 x) {
#if !defined(GGML_USE_HIP) && __CUDA_ARCH__ >= GGML_CUDA_CC_PASCAL || defined(GGML_USE_HIP)
#pragma unroll
for (int offset = width/2; offset > 0; offset >>= 1) {
x = ggml_cuda_hmax2(x, __shfl_xor_sync(0xffffffff, x, offset, width));
}
return x;
#else
GGML_UNUSED(x);
NO_DEVICE_CODE;
#endif // !defined(GGML_USE_HIP) && __CUDA_ARCH__ >= GGML_CUDA_CC_PASCAL || defined(GGML_USE_HIP)
}
#if (defined(CUDART_VERSION) && CUDART_VERSION < CUDART_HMASK) || defined(GGML_USE_HIP) || \
+1 -5
View File
@@ -1875,7 +1875,7 @@ static __global__ void flash_attn_ext_f16(
#endif // defined(AMD_WMMA_AVAILABLE)
#if defined(AMD_MFMA_AVAILABLE)
if (ncols1*ncols2 < 16 || DKQ > 256) {
if (ncols1*ncols2 < 16 || (DKQ > 256 && ncols1*ncols2 < 32)) {
NO_DEVICE_CODE;
return;
}
@@ -2091,26 +2091,22 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml
constexpr bool use_sparse_kernel = false;
fattn_kernel = flash_attn_ext_f16<DKQ, DV, ncols1, ncols2, use_logit_softcap, V_is_K_view, use_sparse_kernel>;
#if !defined(GGML_USE_MUSA)
static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false};
if (!shared_memory_limit_raised[id]) {
CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast<fattn_kernel_ptr_t>(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total));
shared_memory_limit_raised[id] = true;
}
#endif // !defined(GGML_USE_MUSA)
}
} else {
constexpr bool use_logit_softcap = true;
constexpr bool use_sparse_kernel = false;
fattn_kernel = flash_attn_ext_f16<DKQ, DV, ncols1, ncols2, use_logit_softcap, V_is_K_view, use_sparse_kernel>;
#if !defined(GGML_USE_MUSA)
static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false};
if (!shared_memory_limit_raised[id]) {
CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast<fattn_kernel_ptr_t>(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total));
shared_memory_limit_raised[id] = true;
}
#endif // !defined(GGML_USE_MUSA)
}
launch_fattn<DV, ncols1, ncols2>
+13 -13
View File
@@ -19,41 +19,41 @@
} \
static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_nvidia_fp16(const int DKQ, const int DV, const int ncols) {
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 2, 64, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 4, 128, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 8, 256, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 2, 128, 3, 128, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 4, 128, 2, 128, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 8, 128, 3, 128, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 16, 256, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 32, 256, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 2, 64, 2, 64, 64)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 4, 128, 2, 64, 64)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 8, 256, 2, 64, 64)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 8, 256, 3, 128, 64)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 16, 256, 2, 64, 64)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 32, 256, 2, 64, 64)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 2, 64, 2, 64, 72)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 4, 128, 2, 64, 72)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 2, 128, 2, 64, 72)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 4, 128, 3, 128, 72)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 8, 256, 2, 64, 72)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 16, 256, 2, 64, 72)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 32, 256, 2, 64, 72)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 2, 64, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 2, 128, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 4, 128, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 8, 256, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 16, 256, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 32, 256, 2, 64, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 32, 256, 2, 128, 40)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 2, 64, 2, 64, 48)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 4, 128, 2, 64, 48)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 8, 256, 2, 64, 48)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 16, 256, 2, 64, 48)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 32, 256, 2, 64, 48)
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 32, 256, 2, 128, 24)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 2, 64, 2, 64, 56)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 2, 128, 3, 64, 56)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 4, 128, 2, 64, 56)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 8, 256, 2, 64, 56)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 16, 256, 2, 64, 56)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 32, 256, 2, 64, 56)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 8, 256, 3, 64, 112)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 16, 256, 2, 32, 112)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 32, 256, 2, 128, 56)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(128, 128, 2, 64, 2, 64, 64)
GGML_CUDA_FATTN_TILE_CONFIG_CASE(128, 128, 4, 128, 2, 64, 64)
+4 -1
View File
@@ -681,7 +681,7 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
}
// AMD MFMA needs a certain minimum batch size to outscale the tile kernel for large head sizes.
if ((amd_mfma_available(cc) && Q->ne[0] <= 256) && Q->ne[0] != 40 && Q->ne[0] != 72) {
if (amd_mfma_available(cc) && Q->ne[0] != 40 && Q->ne[0] != 72) {
if ((Q->ne[0] <= 64 && Q->ne[1] * gqa_ratio_eff > 8)) {
return BEST_FATTN_KERNEL_MMA_F16;
}
@@ -691,6 +691,9 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
if ((Q->ne[0] <= 256 && Q->ne[1] * gqa_ratio_eff > 64)) {
return BEST_FATTN_KERNEL_MMA_F16;
}
if (Q->ne[0] > 256 && gqa_opt_applies && Q->ne[1] * gqa_ratio_eff > 128) {
return BEST_FATTN_KERNEL_MMA_F16;
}
}
// AMD WMMA is faster than the tile kernel if the wide tiles with high arithmetic intensity can be utilized.
+40 -13
View File
@@ -1,9 +1,10 @@
#include "common.cuh"
#include "convert.cuh"
#include "fwht.cuh"
template <int N>
template <int N, typename T>
__launch_bounds__(4*ggml_cuda_get_physical_warp_size(), 1)
__global__ void fwht_cuda(const float * src, float * dst, const int64_t n_rows, const float scale) {
__global__ void fwht_cuda(const T * src, float * dst, const int64_t n_rows, const float scale) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
const int64_t r = (int64_t) blockIdx.x * blockDim.y + threadIdx.y;
@@ -22,7 +23,7 @@ __global__ void fwht_cuda(const float * src, float * dst, const int64_t n_rows,
ggml_cuda_pdl_sync();
#pragma unroll
for (int i = 0; i < el_w; ++i) {
reg[i] = src[i * warp_size + lane] * scale;
reg[i] = ggml_cuda_cast<float>(src[i * warp_size + lane]) * scale;
}
#pragma unroll
@@ -58,15 +59,12 @@ __global__ void fwht_cuda(const float * src, float * dst, const int64_t n_rows,
}
}
bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
GGML_ASSERT(ggml_are_same_shape(src, dst));
if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
return false;
}
template <typename T>
static bool ggml_cuda_op_fwht_impl(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
const int n = src->ne[0];
const int64_t rows = ggml_nrows(src);
const float * src_d = (const float *) src->data;
const T * src_d = (const T *) src->data;
float * dst_d = (float *) dst->data;
const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size;
@@ -84,18 +82,47 @@ bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src,
switch (n) {
case 64:
ggml_cuda_kernel_launch(fwht_cuda<64>, launch_params, src_d, dst_d, rows, scale);
ggml_cuda_kernel_launch(fwht_cuda<64, T>, launch_params, src_d, dst_d, rows, scale);
return true;
case 128:
ggml_cuda_kernel_launch(fwht_cuda<128>, launch_params, src_d, dst_d, rows, scale);
ggml_cuda_kernel_launch(fwht_cuda<128, T>, launch_params, src_d, dst_d, rows, scale);
return true;
case 256:
ggml_cuda_kernel_launch(fwht_cuda<256>, launch_params, src_d, dst_d, rows, scale);
ggml_cuda_kernel_launch(fwht_cuda<256, T>, launch_params, src_d, dst_d, rows, scale);
return true;
case 512:
ggml_cuda_kernel_launch(fwht_cuda<512>, launch_params, src_d, dst_d, rows, scale);
ggml_cuda_kernel_launch(fwht_cuda<512, T>, launch_params, src_d, dst_d, rows, scale);
return true;
default:
return false;
}
}
bool ggml_cuda_op_mul_mat_use_fwht(const struct ggml_tensor * op) {
const struct ggml_tensor * a = op->src[0];
const struct ggml_tensor * b = op->src[1];
return op->op == GGML_OP_MUL_MAT && ggml_get_op_params_i32(op, 1) == GGML_HINT_SRC0_IS_HADAMARD &&
a->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 &&
(b->type == GGML_TYPE_F32 || b->type == GGML_TYPE_F16) && ggml_is_contiguous(b) && ggml_is_contiguous(op) &&
ggml_are_same_shape(b, op);
}
bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
GGML_ASSERT(ggml_are_same_shape(src, dst));
if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
return false;
}
if (dst->type != GGML_TYPE_F32) {
return false;
}
switch (src->type) {
case GGML_TYPE_F32:
return ggml_cuda_op_fwht_impl<float>(ctx, src, dst);
case GGML_TYPE_F16:
return ggml_cuda_op_fwht_impl<half>(ctx, src, dst);
default:
return false;
}
}
+2
View File
@@ -1,4 +1,6 @@
#include "common.cuh"
bool ggml_cuda_op_mul_mat_use_fwht(const struct ggml_tensor * op);
// Returns whether the Fast Walsh-Hadamard transform could be used.
bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst);
+23 -16
View File
@@ -310,13 +310,9 @@ static ggml_cuda_device_info ggml_cuda_init() {
info.devices[id].smpb = prop.sharedMemPerBlock;
info.devices[id].warp_size = prop.warpSize;
#ifndef GGML_USE_MUSA
int supports_coop_launch = 0;
CUDA_CHECK(cudaDeviceGetAttribute(&supports_coop_launch, cudaDevAttrCooperativeLaunch, physical_id));
info.devices[id].supports_cooperative_launch = !!supports_coop_launch;
#else
info.devices[id].supports_cooperative_launch = false;
#endif // !(GGML_USE_MUSA)
#if defined(GGML_USE_HIP)
info.devices[id].smpbo = prop.sharedMemPerBlock;
@@ -337,8 +333,6 @@ static ggml_cuda_device_info ggml_cuda_init() {
device_vmm ? "yes" : "no", prop.warpSize,
device_vram_mib);
#elif defined(GGML_USE_MUSA)
// FIXME: Ensure compatibility with varying warp sizes across different MUSA archs.
info.devices[id].warp_size = 32;
info.devices[id].smpbo = prop.sharedMemPerBlockOptin;
info.devices[id].cc = GGML_CUDA_CC_OFFSET_MTHREADS + prop.major * 0x100;
info.devices[id].cc += prop.minor * 0x10;
@@ -1824,8 +1818,7 @@ static bool ggml_cuda_should_fuse_mul_mat_vec_q(const ggml_tensor * tensor) {
static void ggml_cuda_mul_mat(ggml_backend_cuda_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
GGML_TENSOR_BINARY_OP_LOCALS
const int32_t hint = ggml_get_op_params_i32(dst, 1);
if (hint == GGML_HINT_SRC0_IS_HADAMARD && ggml_cuda_op_fwht(ctx, src1, dst)) {
if (ggml_cuda_op_mul_mat_use_fwht(dst) && ggml_cuda_op_fwht(ctx, src1, dst)) {
return;
}
@@ -3316,6 +3309,19 @@ static bool ggml_cuda_can_fuse(const struct ggml_cgraph * cgraph,
return true;
}
if (ops.size() == 2 && ops.begin()[0] == GGML_OP_RMS_NORM && ops.begin()[1] == GGML_OP_SCALE) {
const ggml_tensor * rms_norm = cgraph->nodes[node_idx];
const ggml_tensor * scale = cgraph->nodes[node_idx+1];
GGML_ASSERT(rms_norm->src[0]->type == GGML_TYPE_F32);
GGML_ASSERT(rms_norm->type == GGML_TYPE_F32);
float bias;
memcpy(&bias, (const float *) scale->op_params + 1, sizeof(float));
return bias == 0.0f && scale->type == GGML_TYPE_F32;
}
if (ops.size() == 2 && ops.begin()[0] == GGML_OP_SSM_CONV && ops.begin()[1] == GGML_OP_UNARY
&& unary_ops.size() == 1 && unary_ops.begin()[0] == GGML_UNARY_OP_SILU) {
const ggml_tensor * ssm_conv = cgraph->nodes[node_idx];
@@ -4157,6 +4163,11 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
return 1;
}
if (ggml_cuda_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_SCALE }, {})) {
ggml_cuda_op_rms_norm_scale_fused(*cuda_ctx, node, cgraph->nodes[i + 1]);
return 1;
}
if (ggml_cuda_can_fuse(cgraph, i, { GGML_OP_SSM_CONV, GGML_OP_ADD, GGML_OP_UNARY }, { GGML_UNARY_OP_SILU })) {
ggml_cuda_op_ssm_conv(*cuda_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2]);
return 2;
@@ -5200,7 +5211,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
if (a->nb[0] != ggml_element_size(a) || b->nb[0] != ggml_element_size(b)) {
return false; // TODO this could in principle be implemented though currently there is no use case.
}
if (b->type == GGML_TYPE_F16 && a->type != GGML_TYPE_F16) {
if (b->type == GGML_TYPE_F16 && a->type != GGML_TYPE_F16 && !ggml_cuda_op_mul_mat_use_fwht(op)) {
return false;
}
if (op->op == GGML_OP_MUL_MAT_ID && ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
@@ -5476,8 +5487,9 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
if (op->src[3]->ne[0] == 1) {
// Mamba2
// (kernel only supports (d_state == 128 || d_state == 256) && d_head % 16 == 0)
return (op->src[0]->ne[0] == 128 || op->src[0]->ne[0] == 256) && op->src[0]->ne[1] % 16 == 0;
// (kernel only supports (d_state == 96 || d_state == 128 || d_state == 256) && d_head % 16 == 0)
const int64_t d_state = op->src[0]->ne[0];
return (d_state == 96 || d_state == 128 || d_state == 256) && op->src[0]->ne[1] % 16 == 0;
} else {
if (K > 1) {
return false;
@@ -5563,12 +5575,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
case GGML_OP_RWKV_WKV7:
return true;
case GGML_OP_GATED_DELTA_NET:
//TODO: enable once MUSA compiler is solved https://github.com/ggml-org/llama.cpp/pull/19504#issuecomment-4018634327
#ifdef GGML_USE_MUSA
return false;
#else
return true;
#endif // GGML_USE_MUSA
case GGML_OP_DSV4_HC_COMB:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32;
+12 -12
View File
@@ -1631,16 +1631,16 @@ 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_fp4(
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_Q4) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q4);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
int * x_qs = (int *) x_tile;
uint32_t * x_sc = (uint32_t *) (x_qs + 2 * MMQ_TILE_NE_K);
const int txi = threadIdx.x;
constexpr int iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback);
constexpr int iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback, GGML_PREC_Q4);
constexpr int threads_per_row = iter_k / QK_MXFP4; // each thread processes 1 block
constexpr int rows_per_warp = warp_size / threads_per_row;
@@ -1670,12 +1670,12 @@ 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_nvfp4(
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(
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) / 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, prec_src1) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1729,12 +1729,12 @@ 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_nvfp4_nvfp4(
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 iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q4) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q4);
constexpr int iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback, GGML_PREC_Q4);
constexpr int threads_per_row = iter_k / QK_NVFP4; // each thread processes 1 block
constexpr int rows_per_warp = warp_size / threads_per_row;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
uint32_t * x_u32 = (uint32_t *) x_tile;
+4 -4
View File
@@ -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> 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 = GGML_PREC_Q8> 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();
@@ -482,7 +482,7 @@ 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 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 ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
@@ -532,7 +532,7 @@ 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 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 ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
@@ -1180,7 +1180,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<8, 8, int> tile_B;
typedef tile<16, 8, float> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
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 ntx = rows_per_warp / tile_C::I;
constexpr int nfrags = MMQ_TILE_NE_K / tile_A::J;
+61 -4
View File
@@ -5,7 +5,7 @@
#include <cstdint>
static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
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);
@@ -71,9 +71,18 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con
break;
// -----------------------------------------------------------------------
case GGML_TYPE_MXFP4:
// src1 at Q4 uses the native FP4 instructions, which are Blackwell-only
if (prec_src1 == GGML_PREC_Q4) {
mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q4>(ctx, args, stream);
break;
}
mul_mat_q_case<GGML_TYPE_MXFP4>(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);
break;
default:
@@ -82,6 +91,47 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con
}
}
// overrides the src1 precision requested by the graph, "auto" keeps the requested one
static ggml_prec ggml_cuda_mmq_get_prec_env() {
const char * env_c = getenv("GGML_CUDA_MMQ_PREC");
if (env_c == nullptr) {
return GGML_PREC_UNDEFINED;
}
std::string env_cpp = env_c;
for (char & c : env_cpp) {
c = std::tolower(c);
}
if (env_cpp == "q4") {
return GGML_PREC_Q4;
}
if (env_cpp == "q8") {
return GGML_PREC_Q8;
}
if (env_cpp != "auto") {
GGML_LOG_WARN("%s: Unknown value for GGML_CUDA_MMQ_PREC: '%s'. Available: 'q4', 'q8', 'auto'.\n", __func__, env_cpp.c_str());
}
return GGML_PREC_UNDEFINED;
}
// src1 is quantized to Q8_1 unless the FP4 types can use 4-bit activations, in which case they
// default to the native W4A4 instructions on Blackwell.
static ggml_prec ggml_cuda_mmq_get_prec_src1(const ggml_tensor * src0, const ggml_tensor * dst, const int cc) {
static const ggml_prec prec_env = ggml_cuda_mmq_get_prec_env();
ggml_prec prec = prec_env;
if (prec == GGML_PREC_UNDEFINED) {
prec = (ggml_prec) ggml_get_op_params_i32(dst, 3);
}
// Q4 only for the FP4 types on Blackwell
GGML_ASSERT(prec == GGML_PREC_UNDEFINED || prec == GGML_PREC_Q8 || prec == GGML_PREC_Q4);
const bool can_use_q4 = (src0->type == GGML_TYPE_NVFP4 || src0->type == GGML_TYPE_MXFP4) && blackwell_mma_available(cc);
if (prec == GGML_PREC_Q8 || !can_use_q4) {
return GGML_PREC_Q8;
}
return GGML_PREC_Q4;
}
void ggml_cuda_mul_mat_q(
ggml_backend_cuda_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst) {
GGML_ASSERT( src1->type == GGML_TYPE_F32);
@@ -128,7 +178,9 @@ void ggml_cuda_mul_mat_q(
const bool fallback = ne01 % 128 != 0;
const bool use_native_fp4 = blackwell_mma_available(cc) && (src0->type == GGML_TYPE_MXFP4 || src0->type == GGML_TYPE_NVFP4);
const ggml_prec prec_src1 = ggml_cuda_mmq_get_prec_src1(src0, dst, cc);
const bool use_native_fp4 = prec_src1 == GGML_PREC_Q4;
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;
@@ -172,7 +224,7 @@ void ggml_cuda_mul_mat_q(
ne02, ne12, s02, s12, s2,
ne03, ne13, s03, s13, s3,
ne1, ne1};
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream);
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
return;
}
@@ -260,7 +312,7 @@ void ggml_cuda_mul_mat_q(
ne03, ne13, s03, s13, s3,
ne12, ncols_opt};
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream);
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
}
bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t n_experts) {
@@ -389,5 +441,10 @@ bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t
return n_experts > 0;
}
// MUSA: the MMQ kernels compute wrong values on PH1 (MTT S5000).
if (cc == GGML_CUDA_CC_PH1) {
return false;
}
return (!GGML_CUDA_CC_IS_CDNA(cc)) || ne11 < MMQ_DP4A_MAX_BATCH_SIZE;
}
+120 -106
View File
@@ -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) {
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) {
if (GGML_CUDA_CC_IS_AMD(cc)) {
if (GGML_CUDA_CC_IS_GCN(cc)) {
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
@@ -247,6 +247,10 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty
return ggml_cuda_mmq_get_config_rdna2(type, J, fallback);
}
if (blackwell_mma_available(cc)) {
// only src1 at Q4 uses the native FP4 config, higher precisions keep src1 at Q8_1
if (prec_src1 != GGML_PREC_Q4 && (type == GGML_TYPE_NVFP4 || type == GGML_TYPE_MXFP4)) {
return ggml_cuda_mmq_get_config_ampere(type, J, fallback);
}
return ggml_cuda_mmq_get_config_blackwell(type, J, fallback);
}
if (ggml_cuda_highest_compiled_arch(cc) >= GGML_CUDA_CC_VOLTA) {
@@ -258,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) {
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) {
#ifdef GGML_USE_HIP
#ifdef GCN
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
@@ -275,6 +279,10 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t
#endif // CDNA
#else
#ifdef BLACKWELL_MMA_AVAILABLE
// only src1 at Q4 uses the native FP4 config, higher precisions keep src1 at Q8_1
if (prec_src1 != GGML_PREC_Q4 && (type == GGML_TYPE_NVFP4 || type == GGML_TYPE_MXFP4)) {
return ggml_cuda_mmq_get_config_ampere(type, J, fallback);
}
return ggml_cuda_mmq_get_config_blackwell(type, J, fallback);
#elif __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA
return ggml_cuda_mmq_get_config_ampere(type, J, fallback);
@@ -284,79 +292,71 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t
return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback);
#endif // BLACKWELL_MMA_AVAILABLE
#endif // GGML_USE_HIP
GGML_UNUSED_VARS(type, J, fallback);
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 constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).type;
static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).type;
}
static __host__ int ggml_cuda_mmq_get_nthreads(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).nthreads;
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).nthreads;
}
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).nthreads;
}
static __host__ int ggml_cuda_mmq_get_occupancy(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).occupancy;
}
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).occupancy;
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
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 constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).I;
static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
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 constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).J;
static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
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 constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).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) {
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 constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).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) {
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 constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).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) {
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 constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).fallback;
static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).fallback;
}
// ---------------------------------------------------------------------------------------------
@@ -365,8 +365,8 @@ static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const in
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc));
}
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback));
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
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) {
@@ -541,9 +541,9 @@ 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>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_funcs() {
if (!ggml_cuda_mmq_get_config(type, J, fallback).use_mma_data_layout()) {
if (!ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).use_mma_data_layout()) {
switch (type) {
case GGML_TYPE_Q1_0:
return ggml_cuda_mmq_util_funcs(
@@ -690,17 +690,23 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
#ifdef BLACKWELL_MMA_AVAILABLE
switch (type) {
case GGML_TYPE_MXFP4:
return ggml_cuda_mmq_util_funcs(
-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>);
if (prec_src1 == GGML_PREC_Q4) {
return ggml_cuda_mmq_util_funcs(
-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>);
}
break;
case GGML_TYPE_NVFP4:
return ggml_cuda_mmq_util_funcs(
-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>);
if (prec_src1 == GGML_PREC_Q4) {
return ggml_cuda_mmq_util_funcs(
-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>);
}
break;
default:
break;
}
@@ -841,37 +847,37 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
case GGML_TYPE_NVFP4:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
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>);
default:
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
}
}
template <ggml_type type, int J, bool fallback>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
static constexpr __device__ int ggml_cuda_mmq_get_vdr() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback>().vdr;
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vdr;
}
template <ggml_type type, int J, bool fallback>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
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>().load_tiles;
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().load_tiles;
}
template <ggml_type type, int J, bool fallback>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
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>().vec_dot;
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vec_dot;
}
template <ggml_type type, int J, bool fallback>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
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>().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>
template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1 = GGML_PREC_Q8>
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,
@@ -880,25 +886,27 @@ static __device__ __forceinline__ void mul_mat_q_process_tile(
const int tile_x_max_i, const int tile_y_max_j, const int kb0_start, const int kb0_stop) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
constexpr int qk = ggml_cuda_type_traits<type>::qk;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr ggml_cuda_mmq_load_tiles_t load_tiles = ggml_cuda_mmq_get_load_tiles<type, J, fallback>();
constexpr ggml_cuda_mmq_vec_dot_t vec_dot = ggml_cuda_mmq_get_vec_dot<type, J, fallback>();
constexpr ggml_cuda_mmq_write_back_t write_back = ggml_cuda_mmq_get_write_back<type, J, fallback>();
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
constexpr ggml_cuda_mmq_load_tiles_t load_tiles = ggml_cuda_mmq_get_load_tiles<type, J, fallback, prec_src1>();
constexpr ggml_cuda_mmq_vec_dot_t vec_dot = ggml_cuda_mmq_get_vec_dot<type, J, fallback, prec_src1>();
constexpr ggml_cuda_mmq_write_back_t write_back = ggml_cuda_mmq_get_write_back<type, J, fallback, prec_src1>();
extern __shared__ int data_mul_mat_q[];
int * tile_y = data_mul_mat_q + J;
int * tile_x = tile_y + GGML_PAD(J*MMQ_TILE_Y_K, nwarps*warp_size);
#if defined(BLACKWELL_MMA_AVAILABLE)
// FP4 tile stores 8 blocks
constexpr int ne_block = (type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4) ? QK_FP4_MMQ : QK8_1_MMQ;
// FP4 tile stores 8 blocks. src1 above Q4 uses the generic
// Q8_1 tile layout instead of the packed FP4 tile.
constexpr int ne_block = ((type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4) && prec_src1 == GGML_PREC_Q4) ?
QK_FP4_MMQ : QK8_1_MMQ;
#else
constexpr int ne_block = QK8_1_MMQ;
#endif // defined(BLACKWELL_MMA_AVAILABLE)
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback);
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback, prec_src1);
constexpr int blocks_per_iter = ITER_K / qk;
float sum[J*I / (nwarps*warp_size)] = {0.0f};
@@ -950,8 +958,8 @@ 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>
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback), ggml_cuda_mmq_get_occupancy(type, J, fallback))
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
__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,
const int32_t * __restrict__ expert_bounds, float * __restrict__ dst, float * __restrict__ tmp_fixup,
@@ -962,15 +970,15 @@ static __global__ void mul_mat_q(
const uint3 ntx) {
// Skip unused template specializations for faster compilation:
if (ggml_cuda_mmq_get_config(type, J, fallback).type == GGML_TYPE_COUNT) {
if (ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).type == GGML_TYPE_COUNT) {
NO_DEVICE_CODE;
return;
}
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
constexpr int qk = ggml_cuda_type_traits<type>::qk;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
const uint32_t nty = (nrows_x + I - 1) / I; // Number of tiles y
@@ -990,7 +998,7 @@ static __global__ void mul_mat_q(
}
__syncthreads();
if constexpr (!ggml_cuda_mmq_get_stream_k(type, J, fallback)) {
if constexpr (!ggml_cuda_mmq_get_stream_k(type, J, fallback, prec_src1)) {
const uint2 tmp2 = fast_div_modulo(blockIdx.z, nchannels_y);
const int wt = tmp2.x;
const int zt = tmp2.y;
@@ -1053,14 +1061,14 @@ static __global__ void mul_mat_q(
const int offset_x = fastdiv(wt, sample_ratio)*stride_sample_x + fastdiv(zt, channel_ratio)*stride_channel_x + it*I*stride_row_x;
constexpr bool fixup = false;
mul_mat_q_process_tile<type, J, fallback, fixup>
mul_mat_q_process_tile<type, J, fallback, fixup, prec_src1>
(x, offset_x, y + offset_y, ids_dst_shared, dst + offset_dst, tmp_fixup, y_scale_tile,
stride_row_x, ncols_y, stride_col_dst,
tile_x_max_i, tile_y_max_j, 0, blocks_per_ne00.z);
return;
}
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback);
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback, prec_src1);
constexpr int blocks_per_iter = ITER_K / qk;
// kbc == k block continuous, current index in continuous ijk space.
@@ -1147,7 +1155,7 @@ static __global__ void mul_mat_q(
const int offset_x = fastdiv(wt, sample_ratio)*stride_sample_x + fastdiv(zt, channel_ratio)*stride_channel_x + it*I*stride_row_x;
constexpr bool fixup = false; // All but (potentially) the last iterations write their data to dst rather than the fixup buffer.
mul_mat_q_process_tile<type, J, fallback, fixup>
mul_mat_q_process_tile<type, J, fallback, fixup, prec_src1>
(x, offset_x, y + offset_y, ids_dst_shared, dst + offset_dst, tmp_fixup, y_scale_tile,
stride_row_x, ncols_y, stride_col_dst,
tile_x_max_i, tile_y_max_j, kb0_start, kb0_stop);
@@ -1231,24 +1239,24 @@ static __global__ void mul_mat_q(
const int offset_x = fastdiv(wt, sample_ratio)*stride_sample_x + fastdiv(zt, channel_ratio)*stride_channel_x + it*I*stride_row_x;
constexpr bool fixup = true; // Last index writes its data to fixup buffer to avoid data races with other blocks.
mul_mat_q_process_tile<type, J, fallback, fixup>
mul_mat_q_process_tile<type, J, fallback, fixup, prec_src1>
(x, offset_x, y + offset_y, ids_dst_shared, dst + offset_dst, tmp_fixup, y_scale_tile,
stride_row_x, ncols_y, stride_col_dst,
tile_x_max_i, tile_y_max_j, kb0_start, kb0_stop);
}
template <ggml_type type, int J, bool fallback>
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback)/2, 1)
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
__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,
float * __restrict__ tmp_last_tile, const uint3 blocks_per_ne00, const int nrows_x, const int ncols_dst,
const int stride_col_dst, const uint3 nchannels_y, const int stride_channel_dst, const uint3 nsamples_y,
const int stride_sample_dst, const uint3 ntx) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = (ggml_cuda_mmq_get_nthreads(type, J, fallback) / 2) / 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) / 2) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
constexpr int qk = ggml_cuda_type_traits<type>::qk;
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback);
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback, prec_src1);
constexpr int blocks_per_iter = ITER_K / qk;
float sum[J / nwarps] = {0.0f};
@@ -1392,22 +1400,22 @@ 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>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
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;
const int nsm = ggml_cuda_info().devices[id].nsm;
const int warp_size = ggml_cuda_info().devices[id].warp_size;
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc);
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1);
GGML_ASSERT(config.nthreads % warp_size == 0);
const int nwarps = config.nthreads / warp_size;
const int nbytes_shared = mmq_get_nbytes_shared(config, cc);
const dim3 block_dims(warp_size, nwarps, 1);
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, false>), nbytes_shared);
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, true>), nbytes_shared);
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, false, prec_src1>), nbytes_shared);
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, true, prec_src1>), nbytes_shared);
const int nty = (args.nrows_x + config.I - 1) / config.I;
const int ntx = (args.ncols_max + config.J - 1) / config.J;
@@ -1426,8 +1434,8 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
const uint3 channel_ratio_fd = init_fastdiv_values(channel_ratio);
const uint3 sample_ratio_fd = init_fastdiv_values(sample_ratio);
if (!ggml_cuda_mmq_get_stream_k(type, J, fallback, cc)) {
mul_mat_q<type, J, fallback><<<block_nums_xy_tiling, block_dims, nbytes_shared, stream>>>
if (!config.stream_k) {
mul_mat_q<type, J, fallback, prec_src1><<<block_nums_xy_tiling, block_dims, nbytes_shared, stream>>>
(args.x, args.y, args.ids_dst, args.expert_bounds, args.dst, nullptr, args.y_scale,
blocks_per_ne00_fd, args.nrows_x, args.ncols_dst, args.stride_row_x, args.ncols_y, args.nrows_dst,
channel_ratio_fd, nchannels_y_fd, args.stride_channel_x, args.stride_channel_y, args.stride_channel_dst,
@@ -1456,7 +1464,7 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
const dim3 block_nums_fixup(block_nums_stream_k.x, config.I/warp_size, 1);
const dim3 block_dims_fixup(block_dims.x, block_dims.y/2, block_dims.z);
mul_mat_q<type, J, fallback><<<block_nums_stream_k, block_dims, nbytes_shared, stream>>>
mul_mat_q<type, J, fallback, prec_src1><<<block_nums_stream_k, block_dims, nbytes_shared, stream>>>
(args.x, args.y, args.ids_dst, args.expert_bounds, args.dst, tmp_fixup.ptr, args.y_scale,
blocks_per_ne00_fd, args.nrows_x, args.ncols_dst, args.stride_row_x, args.ncols_y, args.nrows_dst,
channel_ratio_fd, nchannels_y_fd, args.stride_channel_x, args.stride_channel_y, args.stride_channel_dst,
@@ -1468,13 +1476,13 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
}
CUDA_CHECK(cudaGetLastError());
mul_mat_q_stream_k_fixup<type, J, fallback><<<block_nums_fixup, block_dims_fixup, 0, stream>>>
mul_mat_q_stream_k_fixup<type, J, fallback, prec_src1><<<block_nums_fixup, block_dims_fixup, 0, stream>>>
(args.ids_dst, args.expert_bounds, args.dst, tmp_fixup.ptr, blocks_per_ne00_fd, args.nrows_x, args.ncols_dst,
args.nrows_dst, nchannels_y_fd, args.stride_channel_dst, nsamples_y_fd, args.stride_sample_dst,
ntx_fd);
}
template <ggml_type type, bool fallback>
template <ggml_type type, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
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;
@@ -1484,7 +1492,7 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
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);
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1);
if (config.type == GGML_TYPE_COUNT) {
continue;
}
@@ -1503,52 +1511,52 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
switch (J_best) {
case 8:
launch_mul_mat_q<type, 8, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 8, fallback, prec_src1>(ctx, args, stream);
break;
case 16:
launch_mul_mat_q<type, 16, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 16, fallback, prec_src1>(ctx, args, stream);
break;
case 24:
launch_mul_mat_q<type, 24, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 24, fallback, prec_src1>(ctx, args, stream);
break;
case 32:
launch_mul_mat_q<type, 32, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 32, fallback, prec_src1>(ctx, args, stream);
break;
case 40:
launch_mul_mat_q<type, 40, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 40, fallback, prec_src1>(ctx, args, stream);
break;
case 48:
launch_mul_mat_q<type, 48, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 48, fallback, prec_src1>(ctx, args, stream);
break;
case 56:
launch_mul_mat_q<type, 56, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 56, fallback, prec_src1>(ctx, args, stream);
break;
case 64:
launch_mul_mat_q<type, 64, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 64, fallback, prec_src1>(ctx, args, stream);
break;
case 72:
launch_mul_mat_q<type, 72, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 72, fallback, prec_src1>(ctx, args, stream);
break;
case 80:
launch_mul_mat_q<type, 80, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 80, fallback, prec_src1>(ctx, args, stream);
break;
case 88:
launch_mul_mat_q<type, 88, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 88, fallback, prec_src1>(ctx, args, stream);
break;
case 96:
launch_mul_mat_q<type, 96, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 96, fallback, prec_src1>(ctx, args, stream);
break;
case 104:
launch_mul_mat_q<type, 104, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 104, fallback, prec_src1>(ctx, args, stream);
break;
case 112:
launch_mul_mat_q<type, 112, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 112, fallback, prec_src1>(ctx, args, stream);
break;
case 120:
launch_mul_mat_q<type, 120, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 120, fallback, prec_src1>(ctx, args, stream);
break;
case 128:
launch_mul_mat_q<type, 128, fallback>(ctx, args, stream);
launch_mul_mat_q<type, 128, fallback, prec_src1>(ctx, args, stream);
break;
default:
fprintf(stderr, "J_best=%d\n", J_best);
@@ -1557,20 +1565,24 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
}
}
template <ggml_type type>
template <ggml_type type, ggml_prec prec_src1 = GGML_PREC_Q8>
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;
mul_mat_q_switch_J<type, fallback>(ctx, args, stream);
mul_mat_q_switch_J<type, fallback, prec_src1>(ctx, args, stream);
} else {
constexpr bool fallback = true;
mul_mat_q_switch_J<type, fallback>(ctx, args, stream);
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) \
// FP4 variant: uses native FP4 MMA instead of keeping src1 at Q8_1.
#define DECL_MMQ_CASE_W4A4(type) \
template void mul_mat_q_case<type, GGML_PREC_Q4>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
extern DECL_MMQ_CASE(GGML_TYPE_Q1_0);
extern DECL_MMQ_CASE(GGML_TYPE_Q2_0);
extern DECL_MMQ_CASE(GGML_TYPE_Q4_0);
@@ -1596,6 +1608,8 @@ extern DECL_MMQ_CASE(GGML_TYPE_IQ4_XS);
// -----------------------------------------
extern DECL_MMQ_CASE(GGML_TYPE_MXFP4);
extern DECL_MMQ_CASE(GGML_TYPE_NVFP4);
extern DECL_MMQ_CASE_W4A4(GGML_TYPE_MXFP4);
extern DECL_MMQ_CASE_W4A4(GGML_TYPE_NVFP4);
// -------------------------------------------------------------------------------------------------------------------------
+44 -11
View File
@@ -73,7 +73,7 @@ static __global__ void group_norm_f32(const float * x, float * dst, const int gr
}
}
template <int block_size, bool do_multiply = false, bool do_add = false>
template <int block_size, bool do_multiply = false, bool do_add = false, bool do_scale = false>
static __global__ void rms_norm_f32(const float * x,
float * dst,
const int ncols,
@@ -96,7 +96,8 @@ static __global__ void rms_norm_f32(const float * x,
const uint3 add_ncols_packed = make_uint3(0, 0, 0),
const uint3 add_nrows_packed = make_uint3(0, 0, 0),
const uint3 add_nchannels_packed = make_uint3(0, 0, 0),
const uint3 add_nsamples_packed = make_uint3(0, 0, 0)) {
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;
@@ -107,6 +108,7 @@ static __global__ void rms_norm_f32(const float * 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;
@@ -148,6 +150,8 @@ static __global__ void rms_norm_f32(const float * x,
} 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];
}
@@ -301,25 +305,27 @@ static void group_norm_f32_cuda(
}
}
template <bool do_scale = false>
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 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);
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>, launch_params,
ggml_cuda_kernel_launch(rms_norm_f32<256, false, false, do_scale>, launch_params,
x, dst, ncols, 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));
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>, 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, 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));
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);
}
}
@@ -367,7 +373,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
x, dst, ncols, 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));
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);
} 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};
@@ -375,7 +381,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
x, dst, ncols, 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));
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);
}
} else {
const uint3 mul_ncols_packed = init_fastdiv_values(mul_ncols);
@@ -394,7 +400,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
x, dst, ncols, 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);
add_nchannels_packed, add_nsamples_packed, 1.0f);
} 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};
@@ -402,7 +408,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
x, dst, ncols, 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);
add_nchannels_packed, add_nsamples_packed, 1.0f);
}
}
}
@@ -499,6 +505,33 @@ void ggml_cuda_op_rms_norm(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
rms_norm_f32_cuda(src0_d, dst_d, ne00, ne01, ne02, ne03, s01, s02, s03, eps, stream);
}
void ggml_cuda_op_rms_norm_scale_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * scale_tensor) {
const ggml_tensor * src0 = dst->src[0];
const float * src0_d = (const float *) src0->data;
float * dst_d = (float *) scale_tensor->data;
cudaStream_t stream = ctx.stream();
GGML_ASSERT(src0->type == GGML_TYPE_F32);
GGML_ASSERT(scale_tensor->type == GGML_TYPE_F32);
GGML_TENSOR_UNARY_OP_LOCALS;
float eps;
memcpy(&eps, dst->op_params, sizeof(float));
GGML_ASSERT(eps >= 0.0f);
float scale;
memcpy(&scale, (const float *) scale_tensor->op_params + 0, sizeof(float));
const size_t ts0 = ggml_type_size(src0->type);
GGML_ASSERT(nb00 == ts0);
const int64_t s01 = nb01 / ts0;
const int64_t s02 = nb02 / ts0;
const int64_t s03 = nb03 / ts0;
rms_norm_f32_cuda<true>(src0_d, dst_d, ne00, ne01, ne02, ne03, s01, s02, s03, eps, stream, scale);
}
void ggml_cuda_op_rms_norm_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * mul_tensor) {
const ggml_tensor * rms_norm_src = (ggml_tensor *) dst->src[0];
float eps = 0.0f;
+2
View File
@@ -8,6 +8,8 @@ void ggml_cuda_op_rms_norm(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
void ggml_cuda_op_rms_norm_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * mul_tensor);
void ggml_cuda_op_rms_norm_scale_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * scale_tensor);
void ggml_cuda_op_rms_norm_fused_add(ggml_backend_cuda_context & ctx,
ggml_tensor * dst,
ggml_tensor * mul_tensor,
+23 -7
View File
@@ -1,6 +1,6 @@
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
#if !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
#define USE_CUB
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
#endif // !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
#ifdef USE_CUB
#include <cub/cub.cuh>
@@ -163,13 +163,19 @@ __global__ void __launch_bounds__(d_state, 1)
const int lane = threadIdx.x % WARP_SIZE;
const int warp_idx = blockIdx.x * c_factor + warp;
ggml_cuda_pdl_sync();
// the last block can have unused warps when n_head*d_head is not a multiple of c_factor
if (warp_idx >= n_head * d_head) {
return;
}
const int head_idx = warp_idx / d_head;
const int head_off = (warp_idx % d_head) * sizeof(float);
const int seq_idx = blockIdx.y;
const int group_off = (head_idx / (n_head / n_group)) * d_state * sizeof(float);
ggml_cuda_pdl_sync();
// TODO: refactor strides to be in elements/floats instead of bytes to be cleaner and consistent with the rest of the codebase
const float * s0_warp = (const float *) ((const char *) src0 + src6[seq_idx] * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
const float * x_warp = (const float *) ((const char *) src1 + (seq_idx * src1_nb3) + (warp_idx * sizeof(float)));
@@ -246,7 +252,17 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
// NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
if (src3_nb1 == sizeof(float)) {
// Mamba-2
if (d_state == 128) {
if (d_state == 96) {
constexpr int threads = 96;
constexpr int num_warps = threads/WARP_SIZE;
const dim3 blocks((n_head * head_dim + (num_warps - 1)) / num_warps, n_seq, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(blocks, threads, 0, stream);
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);
} else if (d_state == 128) {
constexpr int threads = 128;
constexpr int num_warps = threads/WARP_SIZE;
@@ -267,7 +283,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
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);
} else {
GGML_ABORT("doesn't support d_state!=(128 or 256).");
GGML_ABORT("doesn't support d_state!=(96, 128 or 256).");
}
} else {
// Mamba-1
@@ -342,7 +358,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
}
}
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
#if !defined(GGML_USE_HIP)
// ============================================================================
// SSD (State Space Duality) kernels for Mamba-2 prefill (n_tok > SSM_SSD_MIN_TOKENS)
//
@@ -821,7 +837,7 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
GGML_ASSERT(src5->nb[2] <= (size_t)INT_MAX);
GGML_ASSERT(src5->nb[3] <= (size_t)INT_MAX);
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
#if !defined(GGML_USE_HIP)
// Mamba-2 with scalar A per head: use SSD matmul path for long sequences.
// Requires NVIDIA Turing+ otherwise fallback to scan.
const bool is_mamba2 = (src3->nb[1] == sizeof(float));
@@ -50,6 +50,12 @@ SOURCE_MMQ = """// This file has been autogenerated by generate_cu_files.py, do
DECL_MMQ_CASE({type});
"""
TYPES_MMQ_W4A4 = ["GGML_TYPE_MXFP4", "GGML_TYPE_NVFP4"]
SOURCE_MMQ_W4A4 = """
DECL_MMQ_CASE_W4A4({type});
"""
SOURCE_MMF = """// This file has been autogenerated by generate_cu_files.py, do not edit manually.
#include "../mmf.cuh"
@@ -105,6 +111,8 @@ for ncols in [8, 16, 32, 64]:
for type in TYPES_MMQ:
with open(f"mmq-instance-{get_short_name(type)}.cu", "w") as f:
f.write(SOURCE_MMQ.format(type=type))
if type in TYPES_MMQ_W4A4:
f.write(SOURCE_MMQ_W4A4.format(type=type))
for type in range(1, 17):
with open(f"mmf-instance-ncols_{type}.cu", "w") as f:
@@ -3,3 +3,5 @@
#include "../mmq.cuh"
DECL_MMQ_CASE(GGML_TYPE_MXFP4);
DECL_MMQ_CASE_W4A4(GGML_TYPE_MXFP4);
@@ -3,3 +3,5 @@
#include "../mmq.cuh"
DECL_MMQ_CASE(GGML_TYPE_NVFP4);
DECL_MMQ_CASE_W4A4(GGML_TYPE_NVFP4);
+5
View File
@@ -98,7 +98,12 @@ __global__ void topk_moe_cuda(const float * logits,
const float clamp_val,
const float scale_val,
const topk_moe_config config) {
#if defined(GGML_USE_MUSA)
// MUSA: every warp of a partially filled block must reach the barrier below.
const int row = MIN(blockIdx.x * blockDim.y + threadIdx.y, n_rows - 1);
#else
const int row = blockIdx.x * blockDim.y + threadIdx.y;
#endif // defined(GGML_USE_MUSA)
if (row >= n_rows) {
return;
}
+2 -2
View File
@@ -248,11 +248,11 @@
typedef __hip_bfloat16 nv_bfloat16;
typedef __hip_bfloat162 nv_bfloat162;
#if HIP_VERSION >= 60200000
#if HIP_VERSION >= 60300000
#include <hip/hip_fp8.h>
typedef __hip_fp8_e4m3 __nv_fp8_e4m3;
#define FP8_AVAILABLE
#endif // HIP_VERSION >= 60200000
#endif // HIP_VERSION >= 60300000
typedef int8_t int8x4_t __attribute__((ext_vector_type(4)));
typedef uint8_t uint8x4_t __attribute__((ext_vector_type(4)));
+5
View File
@@ -44,6 +44,7 @@
#define cudaDeviceGetPCIBusId musaDeviceGetPCIBusId
#define cudaDeviceProp musaDeviceProp
#define cudaDeviceSynchronize musaDeviceSynchronize
#define cudaDeviceGetAttribute musaDeviceGetAttribute
#define cudaError_t musaError_t
#define cudaErrorMemoryAllocation musaErrorMemoryAllocation
#define cudaErrorPeerAccessAlreadyEnabled musaErrorPeerAccessAlreadyEnabled
@@ -114,6 +115,7 @@
#define cuMemRelease muMemRelease
#define cuMemSetAccess muMemSetAccess
#define cuMemUnmap muMemUnmap
#define cudaDevAttrCooperativeLaunch musaDevAttrCooperativeLaunch
#define cudaFuncAttributeMaxDynamicSharedMemorySize musaFuncAttributeMaxDynamicSharedMemorySize
#define cudaFuncSetAttribute musaFuncSetAttribute
#define cudaMemcpy3DPeerParms musaMemcpy3DPeerParms
@@ -145,6 +147,9 @@
#define cudaStreamCaptureModeRelaxed musaStreamCaptureModeRelaxed
#define cudaStreamBeginCapture musaStreamBeginCapture
#define cudaStreamEndCapture musaStreamEndCapture
#define cudaStreamCaptureStatus musaStreamCaptureStatus
#define cudaStreamCaptureStatusNone musaStreamCaptureStatusNone
#define cudaStreamIsCapturing musaStreamIsCapturing
#define cudaOccupancyMaxActiveBlocksPerMultiprocessor musaOccupancyMaxActiveBlocksPerMultiprocessor
typedef __mt_bfloat16 nv_bfloat16;
-2
View File
@@ -79,9 +79,7 @@ static __global__ void rwkv_wkv7_f32(const int B, const int T, const int C, cons
float state[head_size];
__shared__ float _r[head_size], _w[head_size], _k[head_size], _a[head_size], _b[head_size];
#ifndef GGML_USE_MUSA
#pragma unroll
#endif
for (int i = 0; i < head_size; i++) {
state[i] = s[batch_i * state_size + head_i * head_size * head_size + tid * head_size + i];
}
+2 -1
View File
@@ -23,6 +23,7 @@ include(${HEXAGON_SDK_ROOT}/build/cmake/hexagon_fun.cmake)
include(ExternalProject)
option(GGML_HEXAGON_HTP_DEBUG "ggml-hexagon: enable HTP debug output" OFF)
set(GGML_HEXAGON_HTP_BUILD_TYPE "Release" CACHE STRING "ggml-hexagon: HTP skel build type (Release, RelWithDebInfo, Debug)")
set(GGML_HEXAGON_HTP_CERT "$ENV{HEXAGON_HTP_CERT}" CACHE PATH "ggml-hexagon: enable HTP library signing using certificate")
add_library(htp_iface OBJECT
@@ -64,7 +65,7 @@ function(build_htp_skel V)
SOURCE_DIR ${CMAKE_CURRENT_SOURCE_DIR}/htp BUILD_ALWAYS ON
BUILD_BYPRODUCTS ${CMAKE_CURRENT_BINARY_DIR}/libggml-htp-${V}.so
CMAKE_ARGS
-DCMAKE_BUILD_TYPE=Release
-DCMAKE_BUILD_TYPE=${GGML_HEXAGON_HTP_BUILD_TYPE}
-DCMAKE_TOOLCHAIN_FILE=${CMAKE_CURRENT_SOURCE_DIR}/htp/cmake-toolchain.cmake
-DCMAKE_INSTALL_LIBDIR=${CMAKE_CURRENT_BINARY_DIR}
-DHEXAGON_SDK_ROOT=${HEXAGON_SDK_ROOT}
+507 -78
View File
@@ -63,6 +63,7 @@
#include "htp/rope-ops.h"
#include "htp/ssm-conv.h"
#include "htp/gated-delta-net-ops.h"
#include "htp/argsort-ops.h"
#include "htp_iface.h"
#include "htp-drv.h"
@@ -267,10 +268,10 @@ static inline bool ggml_hexagon_is_repack_type(enum ggml_type type) {
return type == GGML_TYPE_Q4_0 || type == GGML_TYPE_Q4_1 ||
type == GGML_TYPE_Q8_0 || type == GGML_TYPE_IQ4_NL ||
type == GGML_TYPE_MXFP4 || type == GGML_TYPE_Q6_K ||
type == GGML_TYPE_Q4_K;
type == GGML_TYPE_Q4_K || type == GGML_TYPE_Q5_K;
}
// Size of one repacked row in the DSP tiled layout. The Q6_K and Q4_K tiles store uncompressed scales/mins,
// Size of one repacked row in the DSP tiled layout. The Q6_K, Q5_K and Q4_K tiles store uncompressed scales/mins,
// so they are larger than the ggml blocks. For the other repack types the tile has the same size as the ggml blocks.
static inline size_t ggml_hexagon_tiled_row_size(enum ggml_type type, int64_t ne0) {
if (type == GGML_TYPE_Q6_K) {
@@ -279,6 +280,9 @@ static inline size_t ggml_hexagon_tiled_row_size(enum ggml_type type, int64_t ne
if (type == GGML_TYPE_Q4_K) {
return (size_t) (ne0 / 32) * (HTP_MM_WEIGHT_TILE_SIZE_Q4_1 / 32);
}
if (type == GGML_TYPE_Q5_K) {
return (size_t) (ne0 / 32) * (HTP_MM_WEIGHT_TILE_SIZE_Q5_K / 32);
}
return ggml_row_size(type, ne0);
}
@@ -365,6 +369,13 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
struct htp_gdn_kernel_params * kparams
);
static void ggml_hexagon_precompute_sort_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * op,
bool is_top_k,
struct htp_sort_kernel_params * kparams
);
static void ggml_hexagon_precompute_fused_mmnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
@@ -1742,6 +1753,202 @@ static void repack_tiled_q4_K(void * data, const ggml_tensor * t, size_t offset,
GGML_UNUSED(size);
}
// tile layout: see HTP_MM_WEIGHT_TILE_SIZE_Q5_K in htp/matmul-ops.h
static void repack_q5_K_tiled(ggml_tensor * t, const void * data, size_t offset, size_t size) {
GGML_ASSERT(offset == 0);
const block_q5_K * src_matrix = (const block_q5_K *) data;
int64_t ne0 = t->ne[0];
int64_t ne1 = t->ne[1];
int64_t ne2 = t->ne[2];
int64_t ne3 = t->ne[3];
int64_t ne0_padded = hex_round_up(ne0, 32);
int64_t ne1_padded = hex_round_up(ne1, 32);
GGML_ASSERT(ne0 % QK_K == 0);
const int n_col_tiles = ne1_padded / 32;
const int n_k_tiles = ne0_padded / 32;
const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q5_K;
const size_t matrix_size = (size_t) n_col_tiles * n_k_tiles * tile_size;
const int64_t sb_per_row = ne0 / QK_K;
for (int i3 = 0; i3 < ne3; i3++) {
for (int i2 = 0; i2 < ne2; i2++) {
const block_q5_K * src_slice = src_matrix + (i3 * ne2 + i2) * (ne1 * sb_per_row);
uint8_t * matrix_dst = (uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size;
memset(matrix_dst, 0, matrix_size);
for (int64_t r = 0; r < ne1; r++) {
const int ct = (int) (r / 32);
const int row = (int) (r % 32);
const block_q5_K * src_row = src_slice + r * sb_per_row;
for (int kt = 0; kt < n_k_tiles; kt++) {
const int kt_local = kt % 8;
const block_q5_K * b = &src_row[kt / 8];
const float d = GGML_FP16_TO_FP32(b->d);
const float dmin = GGML_FP16_TO_FP32(b->dmin);
uint8_t * tile_dst = matrix_dst + ((size_t) ct * n_k_tiles + kt) * tile_size;
uint8_t * plane = tile_dst + 640;
uint8_t sc, m;
get_scale_min_k4(kt_local, b->scales, &sc, &m);
const float D = d * (float) sc;
const float M = -dmin * (float) m;
const uint8_t * qs_sub = b->qs + (kt_local / 2) * 32;
const int shift = (kt_local & 1) ? 4 : 0;
const uint8_t hbit = (uint8_t) (1 << kt_local);
for (int cp = 0; cp < 16; cp++) {
const uint8_t q0 = (qs_sub[2 * cp + 0] >> shift) & 0x0F;
const uint8_t q1 = (qs_sub[2 * cp + 1] >> shift) & 0x0F;
tile_dst[cp * 32 + row] = (uint8_t) ((q1 << 4) | q0);
const int i = cp / 4;
const int lane = (cp % 4) * 32 + row;
if (b->qh[2 * cp + 0] & hbit) {
plane[lane] |= (uint8_t) (1 << (2 * i));
}
if (b->qh[2 * cp + 1] & hbit) {
plane[lane] |= (uint8_t) (1 << (2 * i + 1));
}
}
ggml_half * scale_dst = (ggml_half *) (tile_dst + 512);
scale_dst[2 * row + 0] = GGML_FP32_TO_FP16(D);
scale_dst[2 * row + 1] = GGML_FP32_TO_FP16(M);
}
}
}
}
GGML_UNUSED(size);
}
// Reverse of repack_q5_K_tiled. Unpacks quants losslessly and normalizes scales/mins. Read-back only.
static void repack_tiled_q5_K(void * data, const ggml_tensor * t, size_t offset, size_t size) {
GGML_ASSERT(offset == 0);
block_q5_K * dst_matrix = (block_q5_K *) data;
int64_t ne0 = t->ne[0];
int64_t ne1 = t->ne[1];
int64_t ne2 = t->ne[2];
int64_t ne3 = t->ne[3];
int64_t ne0_padded = hex_round_up(ne0, 32);
int64_t ne1_padded = hex_round_up(ne1, 32);
GGML_ASSERT(ne0 % QK_K == 0);
const int n_col_tiles = ne1_padded / 32;
const int n_k_tiles = ne0_padded / 32;
const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q5_K;
const size_t matrix_size = (size_t) n_col_tiles * n_k_tiles * tile_size;
const int64_t sb_per_row = ne0 / QK_K;
for (int i3 = 0; i3 < ne3; i3++) {
for (int i2 = 0; i2 < ne2; i2++) {
block_q5_K * dst_slice = dst_matrix + (i3 * ne2 + i2) * (ne1 * sb_per_row);
const uint8_t * matrix_src = (const uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size;
for (int64_t r = 0; r < ne1; r++) {
const int ct = (int) (r / 32);
const int row = (int) (r % 32);
block_q5_K * dst_row = dst_slice + r * sb_per_row;
for (int64_t sb = 0; sb < sb_per_row; sb++) {
block_q5_K * b = &dst_row[sb];
memset(b, 0, sizeof(block_q5_K));
float sub_scales[8];
float sub_mins[8];
for (int kt_local = 0; kt_local < 8; kt_local++) {
const int kt = sb * 8 + kt_local;
const uint8_t * tile_src = matrix_src + ((size_t) ct * n_k_tiles + kt) * tile_size;
const uint8_t * plane = tile_src + 640;
const ggml_half * scale_src = (const ggml_half *) (tile_src + 512);
uint8_t * qs_sub = b->qs + (kt_local / 2) * 32;
const int shift = (kt_local & 1) ? 4 : 0;
const uint8_t hbit = (uint8_t) (1 << kt_local);
for (int cp = 0; cp < 16; cp++) {
const uint8_t val = tile_src[cp * 32 + row];
const uint8_t q0 = val & 0x0F;
const uint8_t q1 = val >> 4;
qs_sub[2 * cp + 0] |= (uint8_t) (q0 << shift);
qs_sub[2 * cp + 1] |= (uint8_t) (q1 << shift);
const int i = cp / 4;
const int lane = (cp % 4) * 32 + row;
if (plane[lane] & (1 << (2 * i))) {
b->qh[2 * cp + 0] |= hbit;
}
if (plane[lane] & (1 << (2 * i + 1))) {
b->qh[2 * cp + 1] |= hbit;
}
}
const float D = GGML_FP16_TO_FP32(scale_src[2 * row + 0]);
const float M = GGML_FP16_TO_FP32(scale_src[2 * row + 1]);
sub_scales[kt_local] = (D > 0.0f) ? D : 0.0f;
sub_mins[kt_local] = (-M > 0.0f) ? -M : 0.0f;
}
float max_scale = 0.0f;
float max_min = 0.0f;
for (int j = 0; j < 8; j++) {
if (sub_scales[j] > max_scale) max_scale = sub_scales[j];
if (sub_mins[j] > max_min) max_min = sub_mins[j];
}
float inv_scale = 0.0f;
if (max_scale > 0.0f) {
b->d = GGML_FP32_TO_FP16(max_scale / 63.0f);
const float d_actual = GGML_FP16_TO_FP32(b->d);
inv_scale = (d_actual > 0.0f) ? (1.0f / d_actual) : 0.0f;
} else {
b->d = GGML_FP32_TO_FP16(0.0f);
}
float inv_min = 0.0f;
if (max_min > 0.0f) {
b->dmin = GGML_FP32_TO_FP16(max_min / 63.0f);
const float dmin_actual = GGML_FP16_TO_FP32(b->dmin);
inv_min = (dmin_actual > 0.0f) ? (1.0f / dmin_actual) : 0.0f;
} else {
b->dmin = GGML_FP32_TO_FP16(0.0f);
}
for (int j = 0; j < 8; j++) {
uint8_t ls = (uint8_t) roundf(inv_scale * sub_scales[j]);
uint8_t lm = (uint8_t) roundf(inv_min * sub_mins[j]);
ls = (std::min)((uint8_t) 63, ls);
lm = (std::min)((uint8_t) 63, lm);
if (j < 4) {
b->scales[j] = ls;
b->scales[j + 4] = lm;
} else {
b->scales[j + 4] = (ls & 0xF) | ((lm & 0xF) << 4);
b->scales[j - 4] |= ((ls >> 4) << 6);
b->scales[j - 0] |= ((lm >> 4) << 6);
}
}
}
}
}
}
GGML_UNUSED(size);
}
static void repack_tensor_tiled(ggml_tensor * tensor, const void * data, size_t size) {
switch (tensor->type) {
case GGML_TYPE_Q4_0:
@@ -1768,6 +1975,10 @@ static void repack_tensor_tiled(ggml_tensor * tensor, const void * data, size_t
repack_mxfp4_tiled(tensor, data, 0, size);
break;
case GGML_TYPE_Q5_K:
repack_q5_K_tiled(tensor, data, 0, size);
break;
case GGML_TYPE_Q6_K:
repack_q6_K_tiled(tensor, data, 0, size);
break;
@@ -1856,6 +2067,12 @@ static void ggml_backend_hexagon_buffer_get_tensor(ggml_backend_buffer_t buffer,
repack_tiled_q4_K(data, tensor, offset, size);
break;
case GGML_TYPE_Q5_K:
GGML_ASSERT(offset == 0);
GGML_ASSERT(offset + size <= ggml_nbytes(tensor));
repack_tiled_q5_K(data, tensor, offset, size);
break;
case GGML_TYPE_Q8_0:
GGML_ASSERT(offset == 0);
GGML_ASSERT(offset + size <= ggml_nbytes(tensor));
@@ -1989,6 +2206,10 @@ static void ggml_backend_hexagon_buffer_get_tensor_2d(ggml_backend_buffer_t buff
repack_tiled_q4_K(temp_buf.data(), tensor, offset, temp_size);
break;
case GGML_TYPE_Q5_K:
repack_tiled_q5_K(temp_buf.data(), tensor, offset, temp_size);
break;
case GGML_TYPE_Q8_0:
repack_tiled_q8_0(temp_buf.data(), tensor, offset, temp_size);
break;
@@ -4494,8 +4715,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
if (is_batched_val && wtype == GGML_TYPE_F16 && group_size > 1) {
// Try grouped path first
const bool use_dma_activation = (src1->nb[1]/sizeof(float) > (size_t)ne00_padded);
if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, use_dma_activation, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
use_grouped = true;
}
}
@@ -4786,13 +5006,48 @@ static bool ggml_hexagon_precompute_binary_params(
const bool is_add_id = op == HTP_OP_ADD_ID;
const bool is_scalar = !is_add_id && src1->ne[0] == 1;
const bool is_transposed = src0->nb[1] < src0_row_size || src1->nb[1] < src1_row_size || dst->nb[1] < dst_row_size;
const bool is_row_bcast = !is_add_id && !is_scalar && !is_transposed &&
src1->ne[0] == src0->ne[0] &&
(src0->ne[1] > 1 || src0->ne[2] > 1 || src0->ne[3] > 1) &&
src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1;
const bool is_same_shape = !is_add_id && !is_scalar && !is_transposed &&
src1->ne[0] == src0->ne[0] &&
(src1->ne[1] == src0->ne[1] || src1->ne[1] == 1) &&
src1->ne[1] == src0->ne[1] &&
(src1->ne[2] == src0->ne[2] || src1->ne[2] == 1) &&
(src1->ne[3] == src0->ne[3] || src1->ne[3] == 1);
const bool is_row_bcast = is_same_shape && src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1;
const bool is_complex = !is_add_id && !is_scalar && !is_same_shape && (src1->ne[0] == src0->ne[0]);
const bool is_complex = !is_add_id && !is_scalar && !is_same_shape && !is_row_bcast && (src1->ne[0] == src0->ne[0]);
const bool is_contig = ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ggml_is_contiguous(dst);
const bool is_scalar_broadcast = !is_add_id && (ggml_nelements(src1) == 1);
if (!is_add_id && is_contig && (ggml_are_same_shape(src0, src1) || is_scalar_broadcast)) {
const uint32_t total_elems = (uint32_t) ggml_nelements(src0);
const uint32_t n_threads = sess->n_threads;
const uint32_t max_chunk_elems = 32768 / elem_size;
const uint32_t min_chunk_elems = 256;
const uint32_t target_chunk_elems = hex_round_up((total_elems + (2 * n_threads) - 1) / (2 * n_threads), 32);
const uint32_t chunk_size = (std::min)(max_chunk_elems, (std::max)(target_chunk_elems, min_chunk_elems));
const uint32_t chunk_bytes = hex_round_up(chunk_size * elem_size, 128);
kparams->kernel_type = HTP_BINARY_KERNEL_CHUNKED;
kparams->n_threads = n_threads;
kparams->rows_per_buffer = 1;
kparams->src0_row_size_aligned = src0_row_size_aligned;
kparams->src1_row_size_aligned = is_scalar_broadcast ? 0 : src1_row_size_aligned;
kparams->dst_row_size_aligned = dst_row_size_aligned;
kparams->src1_size = 0;
kparams->chunk_size = chunk_size;
kparams->chunk_bytes = chunk_bytes;
kparams->is_scalar = is_scalar_broadcast ? 1 : 0;
struct htp_binary_vtcm_layout L;
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
if (L.total_bytes == 0 || L.total_bytes > sess->vtcm_size) {
return false;
}
kparams->vtcm_size = L.total_bytes;
return true;
}
enum htp_binary_kernel_type kernel_type;
size_t src1_size = 0;
@@ -4830,6 +5085,32 @@ static bool ggml_hexagon_precompute_binary_params(
struct htp_binary_vtcm_layout L;
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
if (L.rows_per_buffer == 0 || L.total_bytes > sess->vtcm_size) {
if (!is_add_id && is_contig && (ggml_are_same_shape(src0, src1) || is_scalar_broadcast)) {
const uint32_t total_elems = (uint32_t) ggml_nelements(src0);
const uint32_t n_threads = sess->n_threads;
const uint32_t max_chunk_elems = 32768 / elem_size;
const uint32_t min_chunk_elems = 256;
const uint32_t target_chunk_elems = hex_round_up((total_elems + (2 * n_threads) - 1) / (2 * n_threads), 32);
const uint32_t chunk_size = (std::min)(max_chunk_elems, (std::max)(target_chunk_elems, min_chunk_elems));
const uint32_t chunk_bytes = hex_round_up(chunk_size * elem_size, 128);
kparams->kernel_type = HTP_BINARY_KERNEL_CHUNKED;
kparams->n_threads = n_threads;
kparams->rows_per_buffer = 1;
kparams->src1_row_size_aligned = is_scalar_broadcast ? 0 : src1_row_size_aligned;
kparams->src1_size = 0;
kparams->chunk_size = chunk_size;
kparams->chunk_bytes = chunk_bytes;
kparams->is_scalar = is_scalar_broadcast ? 1 : 0;
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
if (L.total_bytes == 0 || L.total_bytes > sess->vtcm_size) {
return false;
}
kparams->vtcm_size = L.total_bytes;
return true;
}
return false;
}
@@ -4927,49 +5208,51 @@ static void ggml_hexagon_precompute_get_rows_params(
const uint32_t ne12 = src1->ne[2];
const uint32_t nr = ne10 * ne11 * ne12;
const size_t nb01 = src0->nb[1];
const size_t nb1 = dst->nb[1];
const ggml_tensor * src0_base = src0->view_src ? src0->view_src : src0;
const auto * extra = src0_base->buffer && ggml_backend_buffer_is_hexagon(src0_base->buffer) ?
(const ggml_hexagon_tensor_extra *) src0_base->extra : nullptr;
const bool tiled = src0->type == GGML_TYPE_Q4_0 || (extra && (extra->flags & GGML_HEXAGON_TENSOR_REPACK) != 0) ||
sess->needs_repack.count(src0_base) || sess->needs_repack.count(src0);
const bool can_use_dma = (src0->type == dst->type) && (nb01 == nb1);
const bool use_dma = can_use_dma && (ne00 >= 2048);
kparams->use_dma = use_dma ? 1 : 0;
uint32_t chunks_per_row = 1;
uint32_t chunk_size = ne00;
uint32_t total_tasks = nr;
if (use_dma) {
kparams->n_threads = (std::min)((uint32_t)sess->n_threads, nr);
kparams->tasks_per_thread = (nr + kparams->n_threads - 1) / kparams->n_threads;
if (src0->type == dst->type) {
kparams->kernel_type = HTP_GET_ROWS_KERNEL_SAMETYPE;
} else if (tiled) {
kparams->kernel_type = HTP_GET_ROWS_KERNEL_TILED;
} else {
if (src0->type == GGML_TYPE_F32 && nr < sess->n_threads) {
const uint32_t min_chunk_size = 1024;
uint32_t max_chunks = ne00 / min_chunk_size;
if (max_chunks == 0) {
max_chunks = 1;
}
chunks_per_row = (std::min)((sess->n_threads + nr - 1) / nr, max_chunks);
chunk_size = (ne00 + chunks_per_row - 1) / chunks_per_row;
total_tasks = nr * chunks_per_row;
}
kparams->n_threads = (std::min)(total_tasks, (uint32_t)sess->n_threads);
kparams->tasks_per_thread = (total_tasks + kparams->n_threads - 1) / kparams->n_threads;
kparams->kernel_type = HTP_GET_ROWS_KERNEL_FLAT;
}
const uint32_t chunks_per_row = 1;
const uint32_t chunk_size = ne00;
const uint32_t total_tasks = nr;
kparams->n_threads = (std::min)((uint32_t)sess->n_threads, total_tasks);
struct htp_get_rows_vtcm_layout vtcm_layout = {};
while (kparams->n_threads > 0) {
htp_get_rows_vtcm_layout_build(&vtcm_layout, kparams->kernel_type, src0->type, ne00, kparams->n_threads);
if (vtcm_layout.total_bytes <= sess->vtcm_size) {
break;
}
--kparams->n_threads;
}
if (kparams->n_threads == 0 && total_tasks > 0) {
htp_get_rows_vtcm_layout_build(&vtcm_layout, kparams->kernel_type, src0->type, ne00, 1);
}
kparams->vtcm_size = (total_tasks == 0) ? 0 : vtcm_layout.total_bytes;
kparams->tasks_per_thread = kparams->n_threads > 0 ? (total_tasks + kparams->n_threads - 1) / kparams->n_threads : 0;
kparams->chunks_per_row = chunks_per_row;
kparams->chunk_size = chunk_size;
kparams->total_tasks = total_tasks;
kparams->div_ne10 = init_fastdiv_values(ne10);
kparams->div_ne10_ne11 = init_fastdiv_values(ne10 * ne11);
kparams->div_chunks_per_row = init_fastdiv_values(chunks_per_row);
kparams->div_ne02 = init_fastdiv_values(ne02);
kparams->div_ne03 = init_fastdiv_values(ne03);
struct htp_get_rows_vtcm_layout vtcm_layout;
htp_get_rows_vtcm_layout_build(&vtcm_layout, src0->type, ne00, kparams->n_threads);
kparams->vtcm_size = vtcm_layout.total_bytes;
kparams->div_ne10 = ne10 > 0 ? init_fastdiv_values(ne10) : fastdiv_values{0, 0};
kparams->div_ne10_ne11 = (ne10 * ne11) > 0 ? init_fastdiv_values(ne10 * ne11) : fastdiv_values{0, 0};
kparams->div_chunks_per_row = chunks_per_row > 0 ? init_fastdiv_values(chunks_per_row) : fastdiv_values{0, 0};
kparams->div_ne02 = ne02 > 0 ? init_fastdiv_values(ne02) : fastdiv_values{0, 0};
kparams->div_ne03 = ne03 > 0 ? init_fastdiv_values(ne03) : fastdiv_values{0, 0};
}
static void ggml_hexagon_precompute_set_rows_params(
@@ -5267,6 +5550,52 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
if (kparams->n_threads > 0) kparams->div_n_threads = init_fastdiv_values(kparams->n_threads);
}
static void ggml_hexagon_precompute_sort_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * op,
bool is_top_k,
struct htp_sort_kernel_params * kparams
) {
memset(kparams, 0, sizeof(*kparams));
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t ne00 = src0->ne[0];
const uint32_t k = dst->ne[0];
int32_t order = GGML_SORT_ORDER_DESC;
if (!is_top_k) {
order = ((const int32_t *) op->op_params)[0];
}
const uint32_t n_threads_max = sess->n_threads > 0 ? sess->n_threads : 4;
const size_t vtcm_budget = sess->vtcm_size > 0 ? sess->vtcm_size : (8 * 1024 * 1024);
struct htp_sort_vtcm_layout layout;
bool ok = htp_sort_solve_layout(&layout, ne00, total_rows, k, n_threads_max, vtcm_budget, is_top_k);
GGML_ASSERT(ok);
kparams->n_threads = (int32_t) layout.n_threads;
kparams->total_rows = (int32_t) total_rows;
kparams->row_start = 0;
kparams->row_end = (int32_t) total_rows;
kparams->ne00 = (int32_t) ne00;
kparams->k = (int32_t) k;
kparams->order = order;
kparams->is_top_k = is_top_k ? 1 : 0;
kparams->use_dma = 1;
kparams->chunk_elems = (int32_t) layout.chunk_elems;
kparams->n_chunks = (int32_t) layout.n_chunks;
kparams->vtcm_size = (int32_t) layout.total_bytes;
kparams->phase1_slot_size = (int32_t) layout.phase1_slot_size;
kparams->merge_values_off = (int32_t) layout.merge_values_off;
kparams->merge_indices_off = (int32_t) layout.merge_indices_off;
kparams->merge_elems = (int32_t) layout.merge_elems;
kparams->n_slots = (int32_t) layout.n_slots;
}
static void ggml_hexagon_precompute_fused_mmnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0, // W0
@@ -5404,12 +5733,13 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s
case GGML_TYPE_IQ4_NL:
case GGML_TYPE_MXFP4:
case GGML_TYPE_Q4_K:
case GGML_TYPE_Q5_K:
case GGML_TYPE_Q6_K:
if (!ggml_is_contiguous(src0) || ggml_is_permuted(src0)) {
return false;
}
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q5_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
return false;
}
@@ -5487,12 +5817,13 @@ static bool ggml_hexagon_supported_mul_mat_id(const struct ggml_hexagon_session
case GGML_TYPE_IQ4_NL:
case GGML_TYPE_MXFP4:
case GGML_TYPE_Q4_K:
case GGML_TYPE_Q5_K:
case GGML_TYPE_Q6_K:
if (!ggml_is_contiguous(src0) || ggml_is_permuted(src0)) {
return false;
}
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q5_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
return false;
}
@@ -5616,7 +5947,8 @@ static bool ggml_hexagon_supported_unary(const struct ggml_hexagon_session * ses
case GGML_OP_LOG:
break;
case GGML_OP_UNARY:
if (ggml_get_unary_op(op) != GGML_UNARY_OP_ABS) {
if (ggml_get_unary_op(op) != GGML_UNARY_OP_ABS &&
ggml_get_unary_op(op) != GGML_UNARY_OP_STEP) {
return false;
}
break;
@@ -5639,6 +5971,26 @@ static bool ggml_hexagon_supported_unary(const struct ggml_hexagon_session * ses
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_sum(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
if (src0->type != GGML_TYPE_F32) {
return false;
}
if (dst->type != GGML_TYPE_F32) {
return false;
}
if (!ggml_is_contiguous(src0) || !ggml_is_contiguous(dst)) {
return false;
}
return true;
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_sum_rows(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
@@ -5660,6 +6012,26 @@ static bool ggml_hexagon_supported_sum_rows(const struct ggml_hexagon_session *
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_argmax(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
if (src0->type != GGML_TYPE_F32) {
return false;
}
if (dst->type != GGML_TYPE_I32) {
return false;
}
if (!ggml_is_contiguous(src0) || !ggml_is_contiguous(dst)) {
return false;
}
return true;
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_activations(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * src1 = op->src[1];
@@ -5806,19 +6178,36 @@ 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->extra) {
const auto * extra = (const ggml_hexagon_tensor_extra *) src0->extra;
if (extra->flags & GGML_HEXAGON_TENSOR_REPACK) {
if (src0->type == GGML_TYPE_Q4_0 && src0->view_src) {
return false;
}
const ggml_tensor * src0_base = src0->view_src ? src0->view_src : src0;
bool is_repacked = false;
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) {
return false;
}
}
is_repacked = is_repacked || sess->needs_repack.count(src0_base) || sess->needs_repack.count(src0);
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_I32 && src0->ne[0] < 32) {
// View offsets use the raw quantized layout and cannot address a tiled allocation.
if (src0->view_src && is_repacked) {
return false;
}
if (src0->type == GGML_TYPE_Q4_0 && src0->buffer && !is_repacked) {
return false;
}
if (src0->type != dst->type && src0->ne[0] < 32) {
return false;
}
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_F16 &&
src0->type != GGML_TYPE_Q8_0 && src0->type != GGML_TYPE_I32) {
src0->type != GGML_TYPE_Q4_0 && src0->type != GGML_TYPE_Q8_0 && src0->type != GGML_TYPE_I32) {
return false;
}
@@ -5826,13 +6215,30 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
return false;
}
if (src0->type == GGML_TYPE_I32) {
if (dst->type != GGML_TYPE_I32) {
if (src0->type == dst->type) {
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_I32 && src0->type != GGML_TYPE_F16) {
return false;
}
}
else if (dst->type != GGML_TYPE_F32) {
} else if (src0->type == GGML_TYPE_I32) {
return false;
} else if (dst->type != GGML_TYPE_F32) {
return false;
}
// Empty recurrent-state gathers are skipped at execution; do not split the graph for them.
if (ggml_is_empty(op)) {
return true;
}
struct htp_get_rows_kernel_params kparams;
ggml_hexagon_precompute_get_rows_params(sess, src0, src1, dst, &kparams);
if (kparams.n_threads == 0 || (size_t) kparams.vtcm_size > sess->vtcm_size) {
return false;
}
// Q4_0 has no raw fallback. Mark only accepted tensors for repacking.
if (src0->type == GGML_TYPE_Q4_0 && !src0->buffer) {
sess->needs_repack.insert(src0);
}
return true;
@@ -5844,42 +6250,36 @@ static bool ggml_hexagon_supported_argsort(const struct ggml_hexagon_session * s
const struct ggml_tensor * src0 = op->src[0]; // values
const struct ggml_tensor * dst = op; // indices
if (src0->type != GGML_TYPE_F32) {
if (src0->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_I32) {
return false;
}
if (dst->type != GGML_TYPE_I32) {
return false;
}
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads_max = sess->n_threads > 0 ? sess->n_threads : 4;
const size_t vtcm_budget = sess->vtcm_size > 0 ? sess->vtcm_size : (8 * 1024 * 1024);
if (src0->ne[0] > (16*1024)) {
// reject tensors with huge rows for now
struct htp_sort_vtcm_layout layout;
if (!htp_sort_solve_layout(&layout, src0->ne[0], total_rows, dst->ne[0], n_threads_max, vtcm_budget, false)) {
return false;
}
return true;
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_top_k(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0]; // values
const struct ggml_tensor * dst = op; // indices
if (src0->type != GGML_TYPE_F32) {
if (src0->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_I32) {
return false;
}
if (dst->type != GGML_TYPE_I32) {
return false;
}
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads_max = sess->n_threads > 0 ? sess->n_threads : 4;
const size_t vtcm_budget = sess->vtcm_size > 0 ? sess->vtcm_size : (8 * 1024 * 1024);
// Single row uses the threaded chunk+merge path. Multi-row uses one full
// buffer per thread, so it keeps the tighter 64K cap.
const bool single_row = (src0->ne[1] == 1 && src0->ne[2] == 1 && src0->ne[3] == 1);
const int64_t max_ne00 = single_row ? (256*1024) : (64*1024);
if (src0->ne[0] > max_ne00) {
struct htp_sort_vtcm_layout layout;
if (!htp_sort_solve_layout(&layout, src0->ne[0], total_rows, dst->ne[0], n_threads_max, vtcm_budget, true)) {
return false;
}
@@ -6187,9 +6587,11 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
case GGML_OP_CONT: return HTP_OP_CPY;
case GGML_OP_GET_ROWS: return HTP_OP_GET_ROWS;
case GGML_OP_SET_ROWS: return HTP_OP_SET_ROWS;
case GGML_OP_SUM: return HTP_OP_SUM;
case GGML_OP_SUM_ROWS: return HTP_OP_SUM_ROWS;
case GGML_OP_ARGSORT: return HTP_OP_ARGSORT;
case GGML_OP_TOP_K: return HTP_OP_TOP_K;
case GGML_OP_ARGMAX: return HTP_OP_ARGMAX;
case GGML_OP_NORM: return HTP_OP_NORM;
case GGML_OP_L2_NORM: return HTP_OP_L2_NORM;
case GGML_OP_RMS_NORM: return HTP_OP_RMS_NORM;
@@ -6226,6 +6628,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
case GGML_UNARY_OP_TANH: return HTP_OP_UNARY_TANH;
case GGML_UNARY_OP_ABS: return HTP_OP_UNARY_ABS;
case GGML_UNARY_OP_RELU: return HTP_OP_UNARY_RELU;
case GGML_UNARY_OP_STEP: return HTP_OP_UNARY_STEP;
default:
break;
}
@@ -6456,6 +6859,12 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg
node.node,
(struct htp_gdn_kernel_params *)node.kernel_params
);
} else if (node.opcode == HTP_OP_ARGSORT || node.opcode == HTP_OP_TOP_K) {
ggml_hexagon_precompute_sort_params(sess,
node.node,
node.opcode == HTP_OP_TOP_K,
(struct htp_sort_kernel_params *) node.kernel_params
);
}
computed_nodes.push_back(std::move(node));
}
@@ -7094,16 +7503,25 @@ static bool ggml_hexagon_supported_cpy(const struct ggml_hexagon_session * sess,
if (dst->type != GGML_TYPE_F32 && dst->type != GGML_TYPE_F16 &&
dst->type != GGML_TYPE_I32) return false;
const bool is_scalar = (ggml_nelements(src0) == 1 && ggml_nelements(dst) == 1);
const bool sametype = (src0->type == dst->type);
const bool transposed = ggml_is_transposed(src0) || ggml_is_transposed(dst);
const bool sameshape = !transposed && ggml_are_same_shape(src0, dst);
const bool transposed = !is_scalar && (ggml_is_transposed(src0) || ggml_is_transposed(dst));
const bool sameshape = is_scalar || (!transposed && ggml_are_same_shape(src0, dst));
// Same-type copies also support I32.
if (src0->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_I32) {
if (!sameshape) return false;
if (sametype) return true;
if ((src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_I32) ||
(src0->type == GGML_TYPE_I32 && dst->type == GGML_TYPE_F32)) {
return true;
}
return false;
}
// can handle any shape and any same-type (pretty slow if reshaping is required)
if (sametype) return true;
// Type conversion is only supported between F32 and F16.
if (src0->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_I32) return false;
// cannot handle re-shaping and type conversion at the same time
if (!sameshape) return false;
return true;
@@ -7252,10 +7670,18 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
supp = ggml_hexagon_supported_unary(sess, op);
break;
case GGML_OP_SUM:
supp = ggml_hexagon_supported_sum(sess, op);
break;
case GGML_OP_SUM_ROWS:
supp = ggml_hexagon_supported_sum_rows(sess, op);
break;
case GGML_OP_ARGMAX:
supp = ggml_hexagon_supported_argmax(sess, op);
break;
case GGML_OP_SOFT_MAX:
supp = ggml_hexagon_supported_softmax(sess, op);
break;
@@ -7272,6 +7698,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
case GGML_UNARY_OP_GELU:
case GGML_UNARY_OP_GELU_QUICK:
case GGML_UNARY_OP_RELU:
case GGML_UNARY_OP_STEP:
supp = ggml_hexagon_supported_unary(sess, op);
break;
default:
@@ -7758,6 +8185,8 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
"please update hexagon_type to match ggml_type");
static_assert((unsigned int) HTP_TYPE_Q4_K == (unsigned int) GGML_TYPE_Q4_K,
"please update hexagon_type to match ggml_type");
static_assert((unsigned int) HTP_TYPE_Q5_K == (unsigned int) GGML_TYPE_Q5_K,
"please update hexagon_type to match ggml_type");
static_assert((unsigned int) HTP_TYPE_Q6_K == (unsigned int) GGML_TYPE_Q6_K,
"please update hexagon_type to match ggml_type");
+19 -3
View File
@@ -19,6 +19,8 @@
#include "htp/ssm-conv.h"
#include "htp/gated-delta-net-ops.h"
#include "htp/softmax-ops.h"
#include "htp/argsort-ops.h"
#include "htp/get-rows-ops.h"
struct htp_opnode {
ggml_tensor * node { nullptr };
@@ -359,17 +361,31 @@ struct htp_opformat {
} else if (node.opcode == HTP_OP_GATED_DELTA_NET) {
const auto * kparams = (const struct htp_gdn_kernel_params *) node.kernel_params;
const char * path = (kparams->kernel_type == HTP_GDN_KERNEL_HMX_CHUNKED) ? "hmx-chunked" : "hvx-recurrent";
snprintf(str, max_size, "%s-%s vtcm %u",
path,
kparams->kda ? "kda" : "scalar",
snprintf(str, max_size, "%s-%s vtcm %u", path, kparams->kda ? "kda" : "scalar",
(unsigned int) (kparams->vtcm_size ? kparams->vtcm_size : kparams->vtcm_per_thread * kparams->n_threads));
} else if (node.opcode == HTP_OP_MUL || node.opcode == HTP_OP_ADD || node.opcode == HTP_OP_ADD_ID ||
node.opcode == HTP_OP_SUB || node.opcode == HTP_OP_DIV) {
const auto * kparams = (const struct htp_binary_kernel_params *) node.kernel_params;
snprintf(str, max_size, "vtcm %u", (unsigned int) kparams->vtcm_size);
} else if (node.opcode == HTP_OP_ARGSORT || node.opcode == HTP_OP_TOP_K) {
const auto * kparams = (const struct htp_sort_kernel_params *) node.kernel_params;
snprintf(str, max_size, "%s nth %d nchk %d chk %d vtcm %d",
node.opcode == HTP_OP_TOP_K ? "top_k" : "argsort",
(int) kparams->n_threads, (int) kparams->n_chunks,
(int) kparams->chunk_elems, (int) kparams->vtcm_size);
} else if (node.opcode == HTP_OP_GET_ROWS) {
const auto * kparams = (const struct htp_get_rows_kernel_params *) node.kernel_params;
const char * ktype_str = "unknown";
switch (kparams->kernel_type) {
case HTP_GET_ROWS_KERNEL_SAMETYPE: ktype_str = "sametype"; break;
case HTP_GET_ROWS_KERNEL_TILED: ktype_str = "tiled"; break;
case HTP_GET_ROWS_KERNEL_FLAT: ktype_str = "flat"; break;
}
snprintf(str, max_size, "%s%s vtcm %u", ktype_str, kparams->n_threads > 1 ? "-multi" : "", (unsigned int) kparams->vtcm_size);
} else {
snprintf(str, max_size, "----");
}
}
void format(const htp_opnode & node) {
File diff suppressed because it is too large Load Diff
+178
View File
@@ -0,0 +1,178 @@
#ifndef HTP_ARGSORT_OPS_H
#define HTP_ARGSORT_OPS_H
#include <stdint.h>
#include <stddef.h>
#include <stdbool.h>
#include <string.h>
#include "hex-fastdiv.h"
struct htp_sort_kernel_params {
int32_t n_threads;
int32_t total_rows;
int32_t row_start;
int32_t row_end;
int32_t ne00;
int32_t k;
int32_t order; // GGML_SORT_ORDER_ASC (0) or GGML_SORT_ORDER_DESC (1)
int32_t is_top_k; // 1 if TOP_K, 0 if ARGSORT
int32_t use_dma; // 1 if DMA enabled
int32_t chunk_elems;
int32_t n_chunks;
int32_t vtcm_size;
int32_t phase1_slot_size;
int32_t merge_values_off;
int32_t merge_indices_off;
int32_t merge_elems;
int32_t n_slots;
int32_t pad[15];
};
struct htp_sort_vtcm_layout {
size_t total_bytes;
size_t phase1_slot_size;
size_t merge_values_off;
size_t merge_indices_off;
uint32_t chunk_elems;
uint32_t n_chunks;
uint32_t merge_elems;
uint32_t n_threads;
uint32_t n_slots;
};
static inline bool htp_sort_solve_layout(
struct htp_sort_vtcm_layout * layout,
uint32_t ne00,
uint32_t total_rows,
uint32_t k,
uint32_t n_threads_max,
size_t vtcm_budget,
bool is_top_k) {
memset(layout, 0, sizeof(*layout));
if (total_rows > 1) {
uint32_t n_threads = total_rows < n_threads_max ? total_rows : n_threads_max;
uint32_t n_vec = (ne00 + 31) / 32;
uint32_t n_vec_pow2 = 1;
while (n_vec_pow2 < n_vec) n_vec_pow2 <<= 1;
uint32_t ne00_padded = n_vec_pow2 * 32;
size_t values_size = ((ne00_padded * sizeof(float)) + 127) & ~127;
size_t indices_size = ((ne00_padded * sizeof(int32_t)) + 127) & ~127;
size_t spad_per_slot = ((values_size + indices_size) + 255) & ~255;
uint32_t n_slots = 2;
if (spad_per_slot * 2 > vtcm_budget) {
n_slots = 1;
}
size_t spad_per_thread = spad_per_slot * n_slots;
while (n_threads > 1 && (spad_per_thread * n_threads) > vtcm_budget) {
n_threads--;
}
size_t total_bytes = spad_per_thread * n_threads;
if (total_bytes > vtcm_budget) {
return false;
}
layout->total_bytes = total_bytes;
layout->phase1_slot_size = spad_per_slot;
layout->chunk_elems = ne00_padded;
layout->n_chunks = 1;
layout->n_threads = n_threads;
layout->n_slots = n_slots;
return true;
}
uint32_t n_vec = (ne00 + 31) / 32;
uint32_t n_vec_pow2 = 1;
while (n_vec_pow2 < n_vec) n_vec_pow2 <<= 1;
uint32_t n_chunks = 1;
if (ne00 > 1024) {
while (n_chunks * 2 <= n_threads_max && n_chunks * 2 <= n_vec_pow2) {
n_chunks *= 2;
}
}
uint32_t chunk_n_vec = n_vec_pow2 / n_chunks;
uint32_t chunk_elems = chunk_n_vec * 32;
size_t phase1_values_size = ((chunk_elems * sizeof(float)) + 127) & ~127;
size_t phase1_indices_size = ((chunk_elems * sizeof(int32_t)) + 127) & ~127;
size_t phase1_slot_size = ((phase1_values_size + phase1_indices_size) + 255) & ~255;
size_t phase1_total_size = phase1_slot_size * n_chunks;
size_t merge_values_size = 0;
size_t merge_indices_size = 0;
size_t merge_values_off = phase1_total_size;
size_t merge_indices_off = 0;
uint32_t merge_elems = 0;
if (n_chunks > 1) {
if (is_top_k) {
uint32_t local_k = k < chunk_elems ? k : chunk_elems;
uint32_t total_candidates = n_chunks * local_k;
uint32_t merge_n_vec = (total_candidates + 31) / 32;
uint32_t merge_n_vec_pow2 = 1;
while (merge_n_vec_pow2 < merge_n_vec) merge_n_vec_pow2 <<= 1;
merge_elems = merge_n_vec_pow2 * 32;
} else {
merge_elems = n_vec_pow2 * 32;
}
merge_values_size = ((merge_elems * sizeof(float)) + 127) & ~127;
merge_indices_size = ((merge_elems * sizeof(int32_t)) + 127) & ~127;
merge_indices_off = merge_values_off + merge_values_size;
}
size_t total_bytes = phase1_total_size + merge_values_size + merge_indices_size;
if (total_bytes > vtcm_budget && n_chunks > 1) {
n_chunks = 1;
chunk_elems = n_vec_pow2 * 32;
phase1_values_size = ((chunk_elems * sizeof(float)) + 127) & ~127;
phase1_indices_size = ((chunk_elems * sizeof(int32_t)) + 127) & ~127;
phase1_slot_size = ((phase1_values_size + phase1_indices_size) + 255) & ~255;
phase1_total_size = phase1_slot_size;
merge_values_size = 0;
merge_indices_size = 0;
merge_values_off = phase1_total_size;
merge_indices_off = 0;
merge_elems = 0;
total_bytes = phase1_total_size;
}
if (total_bytes > vtcm_budget) {
return false;
}
uint32_t n_slots = 1;
if (n_chunks == 1 && phase1_slot_size * 2 <= vtcm_budget) {
n_slots = 2;
total_bytes = phase1_slot_size * 2;
}
layout->total_bytes = total_bytes;
layout->phase1_slot_size = phase1_slot_size;
layout->merge_values_off = merge_values_off;
layout->merge_indices_off = merge_indices_off;
layout->chunk_elems = chunk_elems;
layout->n_chunks = n_chunks;
layout->merge_elems = merge_elems;
layout->n_threads = n_chunks;
layout->n_slots = n_slots;
return true;
}
#if defined(__cplusplus)
static_assert(sizeof(struct htp_sort_kernel_params) <= 128, "htp_sort_kernel_params is too large for kernel_params blob");
#else
_Static_assert(sizeof(struct htp_sort_kernel_params) <= 128, "htp_sort_kernel_params is too large for kernel_params blob");
#endif
#endif // HTP_ARGSORT_OPS_H
+277
View File
@@ -917,6 +917,154 @@ static void binary_thread_add_id_f32(unsigned int nth, unsigned int ith, void *
dma_queue_flush(dma_q);
}
static inline void hvx_div_scalar_f32(uint8_t * restrict dst, const uint8_t * restrict src, const float val, const uint32_t num_elems) {
hvx_mul_scalar_f32(dst, src, 1.0f / val, num_elems);
}
static inline void hvx_div_scalar_f16(uint8_t * restrict dst, const uint8_t * restrict src, const _Float16 val, const uint32_t num_elems) {
hvx_div_scalar_f16_aa(dst, src, val, num_elems);
}
typedef void (*compute_binary_chunked_t)(
uint8_t * restrict dst,
const uint8_t * restrict src0,
const uint8_t * restrict src1,
const uint32_t num_elems
);
typedef void (*compute_binary_scalar_chunked_f32_t)(
uint8_t * restrict dst,
const uint8_t * restrict src0,
const float val,
const uint32_t num_elems
);
typedef void (*compute_binary_scalar_chunked_f16_t)(
uint8_t * restrict dst,
const uint8_t * restrict src0,
const _Float16 val,
const uint32_t num_elems
);
struct binary_chunked_context {
struct htp_ops_context * octx;
struct htp_binary_vtcm_layout vtcm_layout;
uint8_t * vtcm_base;
uint32_t elem_start;
uint32_t nelem;
uint32_t chunk_size;
uint32_t chunks_per_thread;
uint32_t total_chunks;
bool is_scalar;
float scalar_f32;
_Float16 scalar_f16;
compute_binary_chunked_t compute;
compute_binary_scalar_chunked_f32_t compute_scalar_f32;
compute_binary_scalar_chunked_f16_t compute_scalar_f16;
};
static void binary_thread_chunked(unsigned int nth, unsigned int ith, void * data) {
(void) nth;
struct binary_chunked_context * ctx = (struct binary_chunked_context *) data;
struct htp_ops_context * octx = ctx->octx;
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * src1 = octx->src[1];
const struct htp_tensor * dst = octx->dst;
const uint32_t start_chunk = ctx->chunks_per_thread * ith;
const uint32_t end_chunk = MIN(start_chunk + ctx->chunks_per_thread, ctx->total_chunks);
if (start_chunk >= end_chunk) {
return;
}
const uint32_t src0_type = src0->type;
const size_t elem_size = (src0_type == HTP_TYPE_F32) ? sizeof(float) : sizeof(_Float16);
const uint32_t chunk_size = ctx->chunk_size;
const size_t chunk_bytes = ctx->vtcm_layout.src0_spad_half_size;
FARF(HIGH, "binary-chunked: %d/%d (%u:%u) chunks %u elems %u",
ith, nth, start_chunk, end_chunk, ctx->total_chunks, ctx->nelem);
const struct htp_binary_vtcm_layout * layout = &ctx->vtcm_layout;
uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, ctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
uint8_t * src1_spad_base = ctx->is_scalar ? NULL : (VTCM_LAYOUT_PTR(uint8_t, ctx->vtcm_base, layout->off_src1) + (ith * layout->src1_bytes_per_thread));
uint8_t * dst_spad_base = VTCM_LAYOUT_PTR(uint8_t, ctx->vtcm_base, layout->off_dst) + (ith * layout->dst_bytes_per_thread);
dma_queue * dma_q = octx->ctx->dma[ith];
uint32_t prefetch_chunk = start_chunk;
int spad_idx = 0;
for (int k = 0; k < 2 && prefetch_chunk < end_chunk; k++) {
const uint32_t c_start = ctx->elem_start + prefetch_chunk * chunk_size;
const uint32_t c_end = MIN(c_start + chunk_size, ctx->elem_start + ctx->nelem);
const uint32_t cur_elems = c_end - c_start;
const uint32_t cur_bytes = cur_elems * elem_size;
dma_addr_t s0_curr = src0->data + (size_t) c_start * elem_size;
dma_addr_t d_curr = dst->data + (size_t) c_start * elem_size;
uint8_t * s0_spad = src0_spad_base + spad_idx * chunk_bytes;
uint8_t * d_spad = dst_spad_base + spad_idx * chunk_bytes;
dma_queue_push(dma_q, dma_make_data(d_curr, d_spad), chunk_bytes, chunk_bytes, cur_bytes, 0);
dma_queue_push(dma_q, dma_make_data(s0_spad, s0_curr), chunk_bytes, chunk_bytes, cur_bytes, 1);
if (!ctx->is_scalar) {
dma_addr_t s1_curr = src1->data + (size_t) c_start * elem_size;
uint8_t * s1_spad = src1_spad_base + spad_idx * chunk_bytes;
dma_queue_push(dma_q, dma_make_data(s1_spad, s1_curr), chunk_bytes, chunk_bytes, cur_bytes, 1);
}
prefetch_chunk++;
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t c = start_chunk; c < end_chunk; c++) {
const uint32_t c_start = ctx->elem_start + c * chunk_size;
const uint32_t c_end = MIN(c_start + chunk_size, ctx->elem_start + ctx->nelem);
const uint32_t cur_elems = c_end - c_start;
const uint32_t cur_bytes = cur_elems * elem_size;
uint8_t * d_spad = (uint8_t *) dma_queue_pop(dma_q).src;
uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;
uint8_t * s1_spad = ctx->is_scalar ? NULL : (uint8_t *) dma_queue_pop(dma_q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) c);
if (ctx->is_scalar) {
if (src0_type == HTP_TYPE_F32) {
ctx->compute_scalar_f32(d_spad, s0_spad, ctx->scalar_f32, cur_elems);
} else {
ctx->compute_scalar_f16(d_spad, s0_spad, ctx->scalar_f16, cur_elems);
}
} else {
ctx->compute(d_spad, s0_spad, s1_spad, cur_elems);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) c);
dma_addr_t dst_curr = dst->data + (size_t) c_start * elem_size;
dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), chunk_bytes, chunk_bytes, cur_bytes, 1);
if (prefetch_chunk < end_chunk) {
const uint32_t pc_start = ctx->elem_start + prefetch_chunk * chunk_size;
const uint32_t pc_end = MIN(pc_start + chunk_size, ctx->elem_start + ctx->nelem);
const uint32_t p_elems = pc_end - pc_start;
const uint32_t p_bytes = p_elems * elem_size;
dma_addr_t s0_next = src0->data + (size_t) pc_start * elem_size;
dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), chunk_bytes, chunk_bytes, p_bytes, 1);
if (!ctx->is_scalar) {
dma_addr_t s1_next = src1->data + (size_t) pc_start * elem_size;
dma_queue_push(dma_q, dma_make_data(s1_spad, s1_next), chunk_bytes, chunk_bytes, p_bytes, 1);
}
prefetch_chunk++;
}
}
dma_queue_flush(dma_q);
}
static int execute_op_binary(struct htp_ops_context * octx) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * src1 = octx->src[1];
@@ -932,6 +1080,135 @@ static int execute_op_binary(struct htp_ops_context * octx) {
const size_t src1_row_size = src1->ne[0] * elem_size;
const size_t dst_row_size = dst->ne[0] * elem_size;
if (kparams->kernel_type == HTP_BINARY_KERNEL_CHUNKED) {
const bool is_scalar = kparams->is_scalar || (src1->ne[0] == 1 && src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1);
compute_binary_chunked_t compute = NULL;
compute_binary_scalar_chunked_f32_t compute_scalar_f32 = NULL;
compute_binary_scalar_chunked_f16_t compute_scalar_f16 = NULL;
if (is_scalar) {
if (src0_type == HTP_TYPE_F32) {
switch (octx->op) {
case HTP_OP_ADD: compute_scalar_f32 = hvx_add_scalar_f32; break;
case HTP_OP_SUB: compute_scalar_f32 = hvx_sub_scalar_f32; break;
case HTP_OP_MUL: compute_scalar_f32 = hvx_mul_scalar_f32; break;
case HTP_OP_DIV: compute_scalar_f32 = hvx_div_scalar_f32; break;
default: break;
}
} else if (src0_type == HTP_TYPE_F16) {
switch (octx->op) {
case HTP_OP_ADD: compute_scalar_f16 = hvx_add_scalar_f16; break;
case HTP_OP_SUB: compute_scalar_f16 = hvx_sub_scalar_f16; break;
case HTP_OP_MUL: compute_scalar_f16 = hvx_mul_scalar_f16; break;
case HTP_OP_DIV: compute_scalar_f16 = hvx_div_scalar_f16; break;
default: break;
}
}
if (!compute_scalar_f32 && !compute_scalar_f16) {
return HTP_STATUS_NO_SUPPORT;
}
} else {
if (src0_type == HTP_TYPE_F32) {
switch (octx->op) {
case HTP_OP_ADD: compute = hvx_add_f32; break;
case HTP_OP_SUB: compute = hvx_sub_f32; break;
case HTP_OP_MUL: compute = hvx_mul_f32; break;
case HTP_OP_DIV: compute = hvx_div_f32; break;
default: break;
}
} else if (src0_type == HTP_TYPE_F16) {
switch (octx->op) {
case HTP_OP_ADD: compute = hvx_add_f16; break;
case HTP_OP_SUB: compute = hvx_sub_f16; break;
case HTP_OP_MUL: compute = hvx_mul_f16; break;
case HTP_OP_DIV: compute = hvx_div_f16; break;
default: break;
}
}
if (!compute) {
return HTP_STATUS_NO_SUPPORT;
}
}
const uint32_t total_elems = (uint32_t) (src0->ne[0] * src0->ne[1] * src0->ne[2] * src0->ne[3]);
if (total_elems == 0) {
return HTP_STATUS_OK;
}
uint32_t elem_start = 0;
uint32_t nelem = total_elems;
const uint32_t elems_per_line = (src0_type == HTP_TYPE_F32) ? 32 : 64;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) &&
htp_tensor_mdev_data_aligned(src0) &&
(is_scalar || htp_tensor_mdev_data_aligned(src1)) &&
htp_tensor_is_contiguous(dst, elem_size) &&
htp_tensor_is_contiguous(src0, elem_size) &&
(is_scalar || htp_tensor_is_contiguous(src1, elem_size));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
total_elems, can_split ? elems_per_line : 0,
octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
elem_start = range.start;
nelem = range.count;
}
if (nelem == 0) {
return HTP_STATUS_OK;
}
if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
struct htp_binary_vtcm_layout vtcm_layout;
htp_binary_vtcm_layout_build(&vtcm_layout, kparams, octx->ctx->vtcm_size);
if (vtcm_layout.total_bytes == 0 || vtcm_layout.total_bytes > octx->ctx->vtcm_size) {
return HTP_STATUS_VTCM_TOO_SMALL;
}
const uint32_t chunk_size = kparams->chunk_size > 0 ? kparams->chunk_size : (32768 / elem_size);
const uint32_t total_chunks = (nelem + chunk_size - 1) / chunk_size;
const uint32_t n_threads = (total_chunks >= 2) ? MIN(octx->n_threads, total_chunks) : 1;
const uint32_t chunks_per_thread = (total_chunks + n_threads - 1) / n_threads;
float scalar_f32 = 0.0f;
_Float16 scalar_f16 = 0;
if (is_scalar) {
uint8_t * vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, octx->ctx->vtcm_base, vtcm_layout.off_src1);
dma_queue * dma_q = octx->ctx->dma[0];
dma_queue_push(dma_q, dma_make_data(vtcm_src1, src1->data), 128, 0, elem_size, 1);
dma_queue_pop(dma_q);
if (src0_type == HTP_TYPE_F32) {
scalar_f32 = ((const float *) vtcm_src1)[0];
} else {
scalar_f16 = ((const _Float16 *) vtcm_src1)[0];
}
}
struct binary_chunked_context cctx = {
.octx = octx,
.vtcm_layout = vtcm_layout,
.vtcm_base = (uint8_t *) octx->ctx->vtcm_base,
.elem_start = elem_start,
.nelem = nelem,
.chunk_size = chunk_size,
.chunks_per_thread = chunks_per_thread,
.total_chunks = total_chunks,
.is_scalar = is_scalar,
.scalar_f32 = scalar_f32,
.scalar_f16 = scalar_f16,
.compute = compute,
.compute_scalar_f32 = compute_scalar_f32,
.compute_scalar_f16 = compute_scalar_f16,
};
work_queue_run(octx->ctx->work_queue, binary_thread_chunked, &cctx, n_threads);
return HTP_STATUS_OK;
}
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
+39
View File
@@ -17,6 +17,7 @@ enum htp_binary_kernel_type {
HTP_BINARY_KERNEL_ADD_ID,
HTP_BINARY_KERNEL_COMPLEX,
HTP_BINARY_KERNEL_REPEAT,
HTP_BINARY_KERNEL_CHUNKED,
};
struct htp_binary_kernel_params {
@@ -30,6 +31,10 @@ struct htp_binary_kernel_params {
uint32_t src1_size;
uint32_t vtcm_size;
uint32_t chunk_size;
uint32_t chunk_bytes;
uint32_t is_scalar;
};
#if defined(__cplusplus)
@@ -68,6 +73,40 @@ static inline void htp_binary_vtcm_layout_build(
return;
}
if (kparams->kernel_type == HTP_BINARY_KERNEL_CHUNKED) {
const size_t chunk_bytes = kparams->chunk_bytes;
if (chunk_bytes == 0) {
return;
}
L->src0_bytes_per_thread = 2 * chunk_bytes;
L->src1_bytes_per_thread = kparams->is_scalar ? 0 : (2 * chunk_bytes);
L->dst_bytes_per_thread = 2 * chunk_bytes;
L->src0_spad_half_size = chunk_bytes;
L->src1_spad_half_size = kparams->is_scalar ? 0 : chunk_bytes;
L->dst_spad_half_size = chunk_bytes;
L->rows_per_buffer = 1;
L->src1_size = 0;
const size_t src0_total = n_threads * L->src0_bytes_per_thread;
const size_t src1_total = kparams->is_scalar ? 128 : (n_threads * L->src1_bytes_per_thread);
const size_t dst_total = n_threads * L->dst_bytes_per_thread;
size_t off = 0;
VTCM_LAYOUT_ALLOC(off, off_src0, src0_total);
VTCM_LAYOUT_ALLOC(off, off_src1, src1_total);
VTCM_LAYOUT_ALLOC(off, off_dst, dst_total);
if (off > vtcm_size) {
return;
}
L->total_bytes = off;
return;
}
const size_t spad_row_total = (kparams->kernel_type == HTP_BINARY_KERNEL_SAME_SHAPE)
? 2 * (kparams->src0_row_size_aligned + kparams->src1_row_size_aligned + kparams->dst_row_size_aligned)
: 2 * (kparams->src0_row_size_aligned + kparams->dst_row_size_aligned);
+132 -55
View File
@@ -20,6 +20,7 @@ struct htp_concat_context {
uint32_t nrows;
uint32_t elem_start;
uint32_t nelems;
uint32_t nplanes;
struct fastdiv_values div_ne0;
struct fastdiv_values div_ne1;
struct fastdiv_values div_ne2;
@@ -60,39 +61,47 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t i = start_i; i < end_i; i += block_i) {
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
for (uint32_t p = 0; p < cctx->nplanes; p++) {
const uint32_t i3 = p / dst->ne[2];
const uint32_t i2 = p - i3 * dst->ne[2];
const dma_addr_t src0_plane = src0->data + i2 * src0->nb[2] + i3 * src0->nb[3];
const dma_addr_t src1_plane = src1->data + i2 * src1->nb[2] + i3 * src1->nb[3];
const dma_addr_t dst_plane = dst->data + i2 * dst->nb[2] + i3 * dst->nb[3];
uint32_t src1_width_bytes = current_block_i * sizeof(float);
const dma_addr_t src1_addr = src1->data + i * src1->nb[1];
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
for (uint32_t i = start_i; i < end_i; i += block_i) {
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
uint32_t src0_row_bytes = src0_ne0 * sizeof(float);
const dma_addr_t src0_addr = src0->data + i * src0->nb[1];
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
uint32_t src1_width_bytes = current_block_i * sizeof(float);
const dma_addr_t src1_addr = src1_plane + i * src1->nb[1];
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
dma_queue_pop(dma_q); // src1
uint32_t src0_row_bytes = src0_ne0 * sizeof(float);
const dma_addr_t src0_addr = src0_plane + i * src0->nb[1];
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
dma_queue_pop(dma_q); // src1
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
for (uint32_t j = 0; j < src1_ne0_padded; j += 32) {
#pragma unroll(4)
for (uint32_t ii = 0; ii < current_block_i; ii++) {
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(float));
Q6_vgather_ARMVw(&vtcm_tmp[ii], rt, mu, vv);
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(float);
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
for (uint32_t j = 0; j < src1_ne0_padded; j += 32) {
#pragma unroll(4)
for (uint32_t ii = 0; ii < current_block_i; ii++) {
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(float));
Q6_vgather_ARMVw(&vtcm_tmp[ii], rt, mu, vv);
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(float);
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
dma_queue_pop(dma_q); // src0
const dma_addr_t dst_addr = dst_plane + i * dst->nb[1];
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i);
dma_queue_pop(dma_q);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
dma_queue_pop(dma_q); // src0
const dma_addr_t dst_addr = dst->data + i * dst->nb[1];
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i);
dma_queue_pop(dma_q);
}
}
@@ -131,39 +140,47 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t i = start_i; i < end_i; i += block_i) {
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
for (uint32_t p = 0; p < cctx->nplanes; p++) {
const uint32_t i3 = p / dst->ne[2];
const uint32_t i2 = p - i3 * dst->ne[2];
const dma_addr_t src0_plane = src0->data + i2 * src0->nb[2] + i3 * src0->nb[3];
const dma_addr_t src1_plane = src1->data + i2 * src1->nb[2] + i3 * src1->nb[3];
const dma_addr_t dst_plane = dst->data + i2 * dst->nb[2] + i3 * dst->nb[3];
uint32_t src1_width_bytes = current_block_i * sizeof(__fp16);
const dma_addr_t src1_addr = src1->data + i * src1->nb[1];
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
for (uint32_t i = start_i; i < end_i; i += block_i) {
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
uint32_t src0_row_bytes = src0_ne0 * sizeof(__fp16);
const dma_addr_t src0_addr = src0->data + i * src0->nb[1];
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
uint32_t src1_width_bytes = current_block_i * sizeof(__fp16);
const dma_addr_t src1_addr = src1_plane + i * src1->nb[1];
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
dma_queue_pop(dma_q); // src1
uint32_t src0_row_bytes = src0_ne0 * sizeof(__fp16);
const dma_addr_t src0_addr = src0_plane + i * src0->nb[1];
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
dma_queue_pop(dma_q); // src1
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
for (uint32_t j = 0; j < src1_ne0_padded; j += 64) {
#pragma unroll(4)
for (uint32_t ii = 0; ii < current_block_i; ii++) {
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(__fp16));
Q6_vgather_ARMVh(&vtcm_tmp[ii], rt, mu, vv);
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(__fp16);
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
for (uint32_t j = 0; j < src1_ne0_padded; j += 64) {
#pragma unroll(4)
for (uint32_t ii = 0; ii < current_block_i; ii++) {
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(__fp16));
Q6_vgather_ARMVh(&vtcm_tmp[ii], rt, mu, vv);
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(__fp16);
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
dma_queue_pop(dma_q); // src0
const dma_addr_t dst_addr = dst_plane + i * dst->nb[1];
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i);
dma_queue_pop(dma_q);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
dma_queue_pop(dma_q); // src0
const dma_addr_t dst_addr = dst->data + i * dst->nb[1];
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i);
dma_queue_pop(dma_q);
}
}
@@ -230,6 +247,61 @@ static void concat_generic(unsigned int nth, unsigned int ith, void * data) {
}
}
static bool concat_dim1_contiguous_dma(struct htp_ops_context * octx, int dim, uint32_t type_size) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * src1 = octx->src[1];
const struct htp_tensor * dst = octx->dst;
if (dim != 1 || octx->ctx->mdev.count > 1 ||
(dst->type != HTP_TYPE_F32 && dst->type != HTP_TYPE_F16 && dst->type != HTP_TYPE_I32) ||
src0->type != dst->type || src1->type != dst->type ||
src0->ne[0] != dst->ne[0] || src1->ne[0] != dst->ne[0] ||
src0->ne[2] != dst->ne[2] || src1->ne[2] != dst->ne[2] ||
src0->ne[3] != dst->ne[3] || src1->ne[3] != dst->ne[3] ||
dst->ne[1] != src0->ne[1] + src1->ne[1] ||
!htp_tensor_is_contiguous(src0, type_size) ||
!htp_tensor_is_contiguous(src1, type_size) ||
!htp_tensor_is_contiguous(dst, type_size)) {
return false;
}
const uint32_t src0_row_size = src0->ne[0] * type_size;
const uint32_t src1_row_size = src1->ne[0] * type_size;
// v75+ dma_queue_push() writes a 2D descriptor directly and does not split overflow.
#if __HVX_ARCH__ >= 75
if (src0_row_size > 0xffffffu || src1_row_size > 0xffffffu ||
src0->nb[1] > 0xffffffu || src1->nb[1] > 0xffffffu || dst->nb[1] > 0xffffffu ||
src0->ne[1] > UINT16_MAX || src1->ne[1] > UINT16_MAX) {
return false;
}
#endif
dma_queue * q = octx->ctx->dma[0];
for (uint32_t i3 = 0; i3 < dst->ne[3]; ++i3) {
for (uint32_t i2 = 0; i2 < dst->ne[2]; ++i2) {
dma_addr_t dst_addr = dst->data + i3 * dst->nb[3] + i2 * dst->nb[2];
dma_addr_t src0_addr = src0->data + i3 * src0->nb[3] + i2 * src0->nb[2];
dma_addr_t src1_addr = src1->data + i3 * src1->nb[3] + i2 * src1->nb[2];
if (!dma_queue_push(q, dma_make_data(dst_addr, src0_addr), dst->nb[1], src0->nb[1], src0_row_size, src0->ne[1])) {
dma_queue_flush(q);
dma_queue_push(q, dma_make_data(dst_addr, src0_addr), dst->nb[1], src0->nb[1], src0_row_size, src0->ne[1]);
}
dst_addr += src0->ne[1] * dst->nb[1];
if (!dma_queue_push(q, dma_make_data(dst_addr, src1_addr), dst->nb[1], src1->nb[1], src1_row_size, src1->ne[1])) {
dma_queue_flush(q);
dma_queue_push(q, dma_make_data(dst_addr, src1_addr), dst->nb[1], src1->nb[1], src1_row_size, src1->ne[1]);
}
}
}
dma_queue_flush(q);
return true;
}
int op_concat(struct htp_ops_context * octx) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * src1 = octx->src[1];
@@ -237,12 +309,14 @@ int op_concat(struct htp_ops_context * octx) {
int dim = octx->op_params[0];
bool is_2d = dst->ne[2] == 1 && dst->ne[3] == 1;
const uint32_t type_size = (dst->type == HTP_TYPE_F32 || dst->type == HTP_TYPE_I32) ? 4 : 2;
bool is_src1_transposed = (src1->nb[0] > src1->nb[1]);
bool is_src0_transposed = (src0->nb[0] > src0->nb[1]);
if (concat_dim1_contiguous_dma(octx, dim, type_size)) {
return HTP_STATUS_OK;
}
uint32_t n_threads = octx->n_threads;
struct htp_concat_context cctx;
cctx.octx = octx;
@@ -253,7 +327,9 @@ int op_concat(struct htp_ops_context * octx) {
void (*worker_func)(unsigned int, unsigned int, void *) = concat_generic;
if (dim == 0 && is_2d && is_src1_transposed && !is_src0_transposed) {
const bool rows_ok = src0->nb[0] == type_size && src1->nb[1] == type_size && dst->nb[0] == type_size;
if (dim == 0 && is_src1_transposed && !is_src0_transposed && rows_ok) {
const uint32_t total_rows = dst->ne[1];
const size_t dst_data_row_size = dst->ne[0] * type_size;
uint32_t row_start = 0;
@@ -272,6 +348,7 @@ int op_concat(struct htp_ops_context * octx) {
cctx.row_start = row_start;
cctx.nrows = nrows;
cctx.nplanes = dst->ne[2] * dst->ne[3];
uint32_t block_i = (type_size == 4) ? 32 : 64;
+120
View File
@@ -301,6 +301,86 @@ static void cpy_thread_f32_f16_sameshape(unsigned int nth, unsigned int ith, voi
}
}
static void cpy_thread_i32_f32_sameshape(unsigned int nth, unsigned int ith, void * data) {
struct htp_copy_context * ct = (struct htp_copy_context *) data;
struct htp_ops_context * octx = ct->octx;
cpy_preamble;
const uint32_t dr = ct->src0_nrows_per_thread;
const uint32_t ir0 = ct->row_start + dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);
if (ir0 >= ir1) return;
const uint32_t ne02_ne01 = ne02 * ne01;
uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);
uint32_t rem = ir0 - i03 * ne02_ne01;
uint32_t i02 = fastdiv(rem, &ct->div_ne01);
uint32_t i01 = rem - i02 * ne01;
uint8_t* dst_ptr = (uint8_t*) dst->data + i01*nb1 + i02*nb2 + i03*nb3;
uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
for (uint32_t r = ir0; r < ir1; r++) {
hex_l2fetch(src0_ptr, ne00 * sizeof(float), nb01, 2);
const float * restrict src_row = (const float *) src0_ptr;
int32_t * restrict dst_row = (int32_t *) dst_ptr;
for (uint32_t i = 0; i < ne00; i++) {
dst_row[i] = (int32_t) src_row[i];
}
dst_ptr += nb1;
src0_ptr += nb01;
if (++i01 == ne01) {
i01 = 0;
if (++i02 == ne02) {
i02 = 0;
i03++;
}
dst_ptr = (uint8_t*) dst->data + i02*nb2 + i03*nb3;
src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
}
}
}
static void cpy_thread_f32_i32_sameshape(unsigned int nth, unsigned int ith, void * data) {
struct htp_copy_context * ct = (struct htp_copy_context *) data;
struct htp_ops_context * octx = ct->octx;
cpy_preamble;
const uint32_t dr = ct->src0_nrows_per_thread;
const uint32_t ir0 = ct->row_start + dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);
if (ir0 >= ir1) return;
const uint32_t ne02_ne01 = ne02 * ne01;
uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);
uint32_t rem = ir0 - i03 * ne02_ne01;
uint32_t i02 = fastdiv(rem, &ct->div_ne01);
uint32_t i01 = rem - i02 * ne01;
uint8_t* dst_ptr = (uint8_t*) dst->data + i01*nb1 + i02*nb2 + i03*nb3;
uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
for (uint32_t r = ir0; r < ir1; r++) {
hex_l2fetch(src0_ptr, ne00 * sizeof(int32_t), nb01, 2);
const int32_t * restrict src_row = (const int32_t *) src0_ptr;
float * restrict dst_row = (float *) dst_ptr;
for (uint32_t i = 0; i < ne00; i++) {
dst_row[i] = (float) src_row[i];
}
dst_ptr += nb1;
src0_ptr += nb01;
if (++i01 == ne01) {
i01 = 0;
if (++i02 == ne02) {
i02 = 0;
i03++;
}
dst_ptr = (uint8_t*) dst->data + i02*nb2 + i03*nb3;
src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
}
}
}
static inline void cpy_dma_push_2d_chunked(
dma_queue * dma_q,
dma_addr_t dst,
@@ -366,6 +446,42 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
cpy_preamble;
*use_dma = false;
const uint32_t total_elems_src = ne00 * ne01 * ne02 * ne03;
const uint32_t total_elems_dst = ne0 * ne1 * ne2 * ne3;
if (total_elems_src == 1 && total_elems_dst == 1) {
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F32 && dst->type == HTP_TYPE_I32) {
((int32_t *) dst->data)[0] = (int32_t) (((const float *) src0->data)[0]);
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_I32 && dst->type == HTP_TYPE_F32) {
((float *) dst->data)[0] = (float) (((const int32_t *) src0->data)[0]);
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_I32 && dst->type == HTP_TYPE_I32) {
((int32_t *) dst->data)[0] = ((const int32_t *) src0->data)[0];
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F32 && dst->type == HTP_TYPE_F32) {
((float *) dst->data)[0] = ((const float *) src0->data)[0];
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F16 && dst->type == HTP_TYPE_F16) {
((__fp16 *) dst->data)[0] = ((const __fp16 *) src0->data)[0];
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F32 && dst->type == HTP_TYPE_F16) {
((__fp16 *) dst->data)[0] = (__fp16) (((const float *) src0->data)[0]);
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F16 && dst->type == HTP_TYPE_F32) {
((float *) dst->data)[0] = (float) (((const __fp16 *) src0->data)[0]);
return HTP_STATUS_OK;
}
}
struct htp_copy_context ct;
ct.octx = octx;
@@ -450,6 +566,10 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
copy_fun = cpy_thread_f16_f32_sameshape;
} else if (dst->type == HTP_TYPE_F32 && src0->type == HTP_TYPE_F16) {
copy_fun = cpy_thread_f32_f16_sameshape;
} else if (dst->type == HTP_TYPE_I32 && src0->type == HTP_TYPE_F32) {
copy_fun = cpy_thread_i32_f32_sameshape;
} else if (dst->type == HTP_TYPE_F32 && src0->type == HTP_TYPE_I32) {
copy_fun = cpy_thread_f32_i32_sameshape;
} else {
return HTP_STATUS_NO_SUPPORT;
}
+16
View File
@@ -426,6 +426,22 @@ static inline bool dma_queue_push(dma_queue *q, dma_data ddata, size_t dst_strid
#endif
static inline void dma_sync_read(dma_queue * dma_q, void * dst, dma_addr_t src, size_t bytes) {
const uint32_t b = (uint32_t) bytes;
if (b > 0) {
dma_queue_push(dma_q, dma_make_data(dst, src), b, b, b, 1);
dma_queue_pop(dma_q);
}
}
static inline void dma_sync_write(dma_queue * dma_q, dma_addr_t dst, const void * src, size_t bytes) {
const uint32_t b = (uint32_t) bytes;
if (b > 0) {
dma_queue_push(dma_q, dma_make_data(dst, src), b, b, b, 1);
dma_queue_pop(dma_q);
}
}
#define DMA_CACHE_MAX_SIZE 256U
// Fully assoc LRU cache
+205 -27
View File
@@ -17,6 +17,7 @@
#include "htp-tensor.h"
#include "hvx-utils.h"
#include "hvx-quant.h"
#include "matmul-ops.h"
#include "get-rows-ops.h"
#include "work-queue.h"
@@ -28,6 +29,9 @@ struct get_rows_context {
uint32_t task_start;
uint32_t tasks;
uint32_t tasks_per_thread;
uint32_t tile_size;
uint32_t tile_stride;
bool index_i32;
};
#define get_rows_preamble \
@@ -195,32 +199,191 @@ static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned
dma_queue_flush(dma_q); \
}
#define F32_BYTES(n) ((n) * sizeof(float))
#define F16_BYTES(n) ((n) * sizeof(__fp16))
#define Q8_0_BYTES(n) (((n) / 32) * sizeof(block_q8_0))
GET_ROWS_THREAD_DT_FN(f32, F32_BYTES, int32_t, { if (cur_elems > 0) hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, cur_elems); })
GET_ROWS_THREAD_DT_FN(f32, F32_BYTES, int64_t, { if (cur_elems > 0) hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, cur_elems); })
static __attribute__((noinline)) void compute_get_rows_f16(float * dst_spad, const void * src_spad, uint32_t cur_elems) {
hvx_dequantize_row_f16_f32(dst_spad, src_spad, cur_elems);
}
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int32_t, { hvx_dequantize_row_f16_f32((float *)dst_spad, src_spad, ne00); })
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int64_t, { hvx_dequantize_row_f16_f32((float *)dst_spad, src_spad, ne00); })
static __attribute__((noinline)) void compute_get_rows_q8_0(float * dst_spad, const void * src_spad, uint32_t cur_elems) {
hvx_dequantize_row_q8_0_f32(dst_spad, src_spad, cur_elems);
}
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int32_t, { hvx_dequantize_row_q8_0_f32((float *)dst_spad, src_spad, ne00); })
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int64_t, { hvx_dequantize_row_q8_0_f32((float *)dst_spad, src_spad, ne00); })
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int32_t, { compute_get_rows_f16((float *)dst_spad, src_spad, cur_elems); })
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int64_t, { compute_get_rows_f16((float *)dst_spad, src_spad, cur_elems); })
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int32_t, { compute_get_rows_q8_0((float *)dst_spad, src_spad, cur_elems); })
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) {
const HVX_VectorPred first2 = Q6_Q_vsetq_R(2);
const HVX_VectorPred first4 = Q6_Q_vsetq_R(4);
HVX_Vector vq = Q6_V_vzero();
if (q4) {
const HVX_VectorPred first1 = Q6_Q_vsetq_R(1);
const HVX_VectorPred first3 = Q6_Q_vsetq_R(3);
for (int group = 3; group >= 0; --group) {
const HVX_Vector v = Q6_V_vror_VR(hvx_vmem(tile + group * VLEN), row);
// Four planes contribute bytes at 0, 32, 64 and 96 after rotation.
HVX_Vector packed = Q6_V_vmux_QVV(first1, v, Q6_V_vror_VR(v, 31));
packed = Q6_V_vmux_QVV(first2, packed, Q6_V_vror_VR(v, 62));
packed = Q6_V_vmux_QVV(first3, packed, Q6_V_vror_VR(v, 93));
vq = Q6_V_vmux_QVV(first4, packed, Q6_V_vror_VR(vq, VLEN - 4));
}
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));
} else {
for (int group = 7; group >= 0; --group) {
const HVX_Vector v = Q6_V_vror_VR(hvx_vmem(tile + group * VLEN), 2 * row);
// Two planes contribute halfwords at 0 and 64 after rotation.
const HVX_Vector packed = Q6_V_vmux_QVV(first2, v, Q6_V_vror_VR(v, 62));
vq = Q6_V_vmux_QVV(first4, packed, Q6_V_vror_VR(vq, VLEN - 4));
}
}
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 = 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 *) dst = values;
}
struct get_rows_tiled_task {
dma_addr_t tile_src_base;
dma_addr_t dst_data;
uint32_t row;
};
static inline struct get_rows_tiled_task get_rows_tiled_calc_task(
const struct htp_ops_context * octx,
const struct get_rows_context * grctx,
uint32_t i,
uint32_t n_k_tiles,
uint32_t tile_size
) {
const struct htp_get_rows_kernel_params * kparams = grctx->kparams;
get_rows_preamble;
const uint32_t i12 = fastdiv(i, &kparams->div_ne10_ne11);
const uint32_t rem = i - i12 * ne11 * ne10;
const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);
const uint32_t i10 = rem - i11 * ne10;
const dma_addr_t src1_data = octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12;
const uint32_t i01 = grctx->index_i32 ? *(const int32_t *)(uintptr_t) src1_data : (uint32_t) *(const int64_t *)(uintptr_t) src1_data;
assert(i01 < ne01);
const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);
const uint32_t i02 = i11 - q02 * ne02;
const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);
const uint32_t i03 = i12 - q03 * ne03;
const uint32_t column_tile = i01 / HTP_MM_HMX_TILE_N_ROWS;
const uint32_t row = i01 % HTP_MM_HMX_TILE_N_ROWS;
const dma_addr_t matrix = octx->src[0]->data + i02*nb02 + i03*nb03;
struct get_rows_tiled_task task;
task.tile_src_base = matrix + (column_tile * n_k_tiles) * tile_size;
task.dst_data = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3;
task.row = row;
return task;
}
static void get_rows_thread_tiled(unsigned int nth, unsigned int ith, void * data) {
struct get_rows_context * grctx = (struct get_rows_context *) data;
struct htp_ops_context * octx = grctx->octx;
const struct htp_get_rows_kernel_params * kparams = grctx->kparams;
get_rows_preamble;
const uint32_t dr = grctx->tasks_per_thread;
const uint32_t ir0 = grctx->task_start + dr * ith;
if (ir0 >= grctx->task_start + grctx->tasks) {
return;
}
const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);
const uint32_t n_k_tiles = ne00 / HTP_MM_HMX_TILE_N_COLS;
const struct htp_get_rows_vtcm_layout * vtcm_layout = &grctx->vtcm_layout;
uint8_t * src_spad_base = grctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;
uint8_t * dst_spad_base = grctx->vtcm_base + vtcm_layout->off_dst + ith * vtcm_layout->dst_bytes_per_thread;
dma_queue * dma_q = octx->ctx->dma[ith];
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
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);
for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, ++spad_idx) {
const uint32_t i = ir0 + step;
struct get_rows_tiled_task task = get_rows_tiled_calc_task(octx, grctx, i, n_k_tiles, tile_size);
// Dummy writeback to prime the queue with dst descriptor
dma_queue_push(dma_q,
dma_make_data(task.dst_data, dst_spad_base + spad_idx * vtcm_layout->dst_spad_half_size),
dst_bytes, vtcm_layout->dst_spad_half_size, dst_bytes, 0);
// Prefetch row tiles
dma_queue_push(dma_q,
dma_make_data(src_spad_base + spad_idx * vtcm_layout->src0_spad_half_size, task.tile_src_base),
tile_stride, tile_size, tile_size, n_k_tiles);
}
for (uint32_t step = 0; step < ir1 - ir0; ++step) {
const uint32_t i = ir0 + step;
float * dst_spad = (float *) dma_queue_pop(dma_q).src;
uint8_t * src_spad = (uint8_t *) dma_queue_pop(dma_q).dst;
struct get_rows_tiled_task task = get_rows_tiled_calc_task(octx, grctx, i, n_k_tiles, tile_size);
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
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);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
// Real writeback of dst_spad
dma_queue_push(dma_q,
dma_make_data(task.dst_data, dst_spad),
dst_bytes, vtcm_layout->dst_spad_half_size, dst_bytes, 1);
const uint32_t next_step = step + 2;
if (next_step < ir1 - ir0) {
const uint32_t ni = ir0 + next_step;
struct get_rows_tiled_task next_task = get_rows_tiled_calc_task(octx, grctx, ni, n_k_tiles, tile_size);
dma_queue_push(dma_q,
dma_make_data(src_spad, next_task.tile_src_base),
tile_stride, tile_size, tile_size, n_k_tiles);
}
}
dma_queue_flush(dma_q);
}
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_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_Q8_0 &&
octx->src[0]->type != HTP_TYPE_I32) {
return HTP_STATUS_NO_SUPPORT;
}
if ((octx->src[0]->type == HTP_TYPE_I32 && octx->dst->type != HTP_TYPE_I32) ||
(octx->src[0]->type != HTP_TYPE_I32 && octx->dst->type != HTP_TYPE_F32)) {
return HTP_STATUS_NO_SUPPORT;
if (kparams->kernel_type == HTP_GET_ROWS_KERNEL_SAMETYPE) {
if (octx->src[0]->type != octx->dst->type) {
return HTP_STATUS_NO_SUPPORT;
}
} else {
if (octx->dst->type != HTP_TYPE_F32) {
return HTP_STATUS_NO_SUPPORT;
}
}
if (octx->src[1]->type != HTP_TYPE_I32 && octx->src[1]->type != HTP_TYPE_I64) {
@@ -262,33 +425,48 @@ int op_get_rows(struct htp_ops_context * octx) {
grctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base;
grctx.task_start = task_start;
grctx.tasks = tasks;
grctx.tasks_per_thread = fastdiv(tasks + n_threads - 1, &octx->n_threads_div);
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_stride = (grctx.tile_size + 127) & ~127;
grctx.index_i32 = octx->src[1]->type == HTP_TYPE_I32;
const uint32_t ne00 = octx->src[0]->ne[0];
htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, octx->src[0]->type, ne00, n_threads);
htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, kparams->kernel_type, octx->src[0]->type, ne00, n_threads);
if (grctx.vtcm_layout.total_bytes > octx->ctx->vtcm_size) {
FARF(ERROR, "get-rows: VTCM reservation %zu is too small, needed %zu\n",
octx->ctx->vtcm_size, grctx.vtcm_layout.total_bytes);
return HTP_STATUS_INVAL_PARAMS;
}
const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32);
work_queue_func_t q_func = NULL;
if (kparams->use_dma) {
q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t);
} else {
switch (octx->src[0]->type) {
case HTP_TYPE_F32: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f32_int32_t : get_rows_thread_f32_int64_t); break;
case HTP_TYPE_F16: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f16_int32_t : get_rows_thread_f16_int64_t); break;
case HTP_TYPE_Q8_0: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_q8_0_int32_t : get_rows_thread_q8_0_int64_t); break;
case HTP_TYPE_I32: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t); break;
default: return HTP_STATUS_NO_SUPPORT;
}
switch (kparams->kernel_type) {
case HTP_GET_ROWS_KERNEL_SAMETYPE:
q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t);
break;
case HTP_GET_ROWS_KERNEL_TILED:
q_func = get_rows_thread_tiled;
break;
case HTP_GET_ROWS_KERNEL_FLAT:
switch (octx->src[0]->type) {
case HTP_TYPE_F16: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f16_int32_t : get_rows_thread_f16_int64_t); break;
case HTP_TYPE_Q8_0: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_q8_0_int32_t : get_rows_thread_q8_0_int64_t); break;
default: return HTP_STATUS_NO_SUPPORT;
}
break;
default:
return HTP_STATUS_NO_SUPPORT;
}
FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu use-dma %d n-threads %d\n",
FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu kernel-type %d n-threads %d\n",
octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3],
octx->src[1]->ne[0], octx->src[1]->ne[1], octx->src[1]->ne[2], octx->src[1]->ne[3],
octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3],
grctx.vtcm_layout.src0_bytes_per_thread * n_threads,
grctx.vtcm_layout.dst_bytes_per_thread * n_threads,
kparams->use_dma, n_threads);
kparams->kernel_type, n_threads);
work_queue_run(octx->ctx->work_queue, q_func, &grctx, n_threads);
return HTP_STATUS_OK;
+34 -6
View File
@@ -1,11 +1,21 @@
#ifndef HTP_GET_ROWS_OPS_H
#define HTP_GET_ROWS_OPS_H
#include <stdbool.h>
#include <string.h>
#include "hex-fastdiv.h"
#include "matmul-ops.h"
enum htp_get_rows_kernel_type {
HTP_GET_ROWS_KERNEL_SAMETYPE = 0,
HTP_GET_ROWS_KERNEL_TILED,
HTP_GET_ROWS_KERNEL_FLAT,
};
struct htp_get_rows_kernel_params {
int32_t n_threads;
int32_t use_dma;
int32_t kernel_type;
int32_t chunks_per_row;
int32_t chunk_size;
int32_t total_tasks;
@@ -34,19 +44,37 @@ struct htp_get_rows_vtcm_layout {
static inline void htp_get_rows_vtcm_layout_build(
struct htp_get_rows_vtcm_layout * vtcm_layout,
int kernel_type,
int type,
uint32_t ne00,
uint32_t n_threads) {
if (kernel_type == HTP_GET_ROWS_KERNEL_SAMETYPE) {
memset(vtcm_layout, 0, sizeof(*vtcm_layout));
return;
}
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_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;
vtcm_layout->src0_spad_half_size = (row_tiles_size + 255) & ~255;
vtcm_layout->dst_spad_half_size = (ne00 * sizeof(float) + 255) & ~255;
vtcm_layout->src0_bytes_per_thread = 2 * vtcm_layout->src0_spad_half_size;
vtcm_layout->dst_bytes_per_thread = 2 * vtcm_layout->dst_spad_half_size;
vtcm_layout->off_src0 = 0;
vtcm_layout->off_dst = vtcm_layout->src0_bytes_per_thread * n_threads;
vtcm_layout->total_bytes = vtcm_layout->off_dst + vtcm_layout->dst_bytes_per_thread * n_threads;
return;
}
uint32_t src0_row_size = 0;
switch (type) {
case 0: // HTP_TYPE_F32
src0_row_size = ne00 * 4;
break;
case 1: // HTP_TYPE_F16
case HTP_TYPE_F16:
src0_row_size = ne00 * 2;
break;
case 8: // HTP_TYPE_Q8_0
case HTP_TYPE_Q8_0:
src0_row_size = (ne00 / 32) * 34;
break;
default:
+105 -92
View File
@@ -506,6 +506,111 @@ static void dequantize_tiled_weight_to_fp16_task_q8_0(
}
}
static void dequantize_tiled_weight_to_fp16_task_q5_k(
const tiled_dequantize_state_t *state,
uint32_t start_tile, uint32_t end_tile) {
const HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
for (uint32_t t = start_tile; t < end_tile; t++) {
const uint8_t * tile_src = state->src + t * state->aligned_tile_size;
__fp16 * dst_ptr = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
HVX_Vector vscale_offset = hvx_vmem(tile_src + 512);
HVX_VectorPair dm_deal = Q6_W_vdeal_VVR(vscale_offset, vscale_offset, -2);
HVX_Vector vd = Q6_V_lo_W(dm_deal);
HVX_Vector vm = Q6_V_hi_W(dm_deal);
HVX_Vector v_scale_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(vd, vd, -2));
HVX_Vector v_offset_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(vm, vm, -2));
// Load all 4 groups in parallel
HVX_Vector vq0 = hvx_vmem(tile_src + 0 * 128);
HVX_Vector vq1 = hvx_vmem(tile_src + 1 * 128);
HVX_Vector vq2 = hvx_vmem(tile_src + 2 * 128);
HVX_Vector vq3 = hvx_vmem(tile_src + 3 * 128);
// Nibble extraction
HVX_Vector v_lo0 = Q6_V_vand_VV(vq0, mask_h4);
HVX_Vector v_hi0 = Q6_Vub_vlsr_VubR(vq0, 4);
HVX_Vector v_lo1 = Q6_V_vand_VV(vq1, mask_h4);
HVX_Vector v_hi1 = Q6_Vub_vlsr_VubR(vq1, 4);
HVX_Vector v_lo2 = Q6_V_vand_VV(vq2, mask_h4);
HVX_Vector v_hi2 = Q6_Vub_vlsr_VubR(vq2, 4);
HVX_Vector v_lo3 = Q6_V_vand_VV(vq3, mask_h4);
HVX_Vector v_hi3 = Q6_Vub_vlsr_VubR(vq3, 4);
// Q5_K: OR in the 5th bit from the plane
HVX_Vector v_plane = hvx_vmem(tile_src + 640);
v_lo0 = hvx_q5k_or_hibit(v_lo0, v_plane, 0);
v_hi0 = hvx_q5k_or_hibit(v_hi0, v_plane, 1);
v_lo1 = hvx_q5k_or_hibit(v_lo1, v_plane, 2);
v_hi1 = hvx_q5k_or_hibit(v_hi1, v_plane, 3);
v_lo2 = hvx_q5k_or_hibit(v_lo2, v_plane, 4);
v_hi2 = hvx_q5k_or_hibit(v_hi2, v_plane, 5);
v_lo3 = hvx_q5k_or_hibit(v_lo3, v_plane, 6);
v_hi3 = hvx_q5k_or_hibit(v_hi3, v_plane, 7);
// Shuffling
HVX_VectorPair vp_shuf0 = Q6_W_vshuff_VVR(v_hi0, v_lo0, -1);
HVX_VectorPair vp_shuf1 = Q6_W_vshuff_VVR(v_hi1, v_lo1, -1);
HVX_VectorPair vp_shuf2 = Q6_W_vshuff_VVR(v_hi2, v_lo2, -1);
HVX_VectorPair vp_shuf3 = Q6_W_vshuff_VVR(v_hi3, v_lo3, -1);
// Unpack to 16-bit
HVX_VectorPair vp_int16_lo0 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf0));
HVX_VectorPair vp_int16_hi0 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf0));
HVX_VectorPair vp_int16_lo1 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf1));
HVX_VectorPair vp_int16_hi1 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf1));
HVX_VectorPair vp_int16_lo2 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf2));
HVX_VectorPair vp_int16_hi2 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf2));
HVX_VectorPair vp_int16_lo3 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf3));
HVX_VectorPair vp_int16_hi3 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf3));
// Convert, multiply, add offset
HVX_Vector v_grp0_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo0)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp0_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo0)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp0_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi0)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp0_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi0)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp1_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo1)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp1_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo1)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp1_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi1)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp1_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi1)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp2_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo2)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp2_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo2)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp2_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi2)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp2_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi2)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp3_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo3)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp3_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo3)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp3_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi3)), v_scale_duplicated), v_offset_duplicated));
HVX_Vector v_grp3_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi3)), v_scale_duplicated), v_offset_duplicated));
// Parallel Stores
hvx_vmem(dst_ptr + 0 * 64) = v_grp0_0;
hvx_vmem(dst_ptr + 1 * 64) = v_grp0_1;
hvx_vmem(dst_ptr + 2 * 64) = v_grp0_2;
hvx_vmem(dst_ptr + 3 * 64) = v_grp0_3;
hvx_vmem(dst_ptr + 4 * 64) = v_grp1_0;
hvx_vmem(dst_ptr + 5 * 64) = v_grp1_1;
hvx_vmem(dst_ptr + 6 * 64) = v_grp1_2;
hvx_vmem(dst_ptr + 7 * 64) = v_grp1_3;
hvx_vmem(dst_ptr + 8 * 64) = v_grp2_0;
hvx_vmem(dst_ptr + 9 * 64) = v_grp2_1;
hvx_vmem(dst_ptr + 10 * 64) = v_grp2_2;
hvx_vmem(dst_ptr + 11 * 64) = v_grp2_3;
hvx_vmem(dst_ptr + 12 * 64) = v_grp3_0;
hvx_vmem(dst_ptr + 13 * 64) = v_grp3_1;
hvx_vmem(dst_ptr + 14 * 64) = v_grp3_2;
hvx_vmem(dst_ptr + 15 * 64) = v_grp3_3;
}
}
// Q6_K stores 6-bit weights and one fp16 scale per 16 k, see HTP_MM_WEIGHT_TILE_SIZE_Q6_K.
// A k-group holds 4 k per row, the HMX tile holds 2, so each group is dealt into two tiles.
static void dequantize_tiled_weight_to_fp16_task_q6_k(
@@ -914,98 +1019,6 @@ typedef struct {
// activations : fp32 -> fp16
static void transfer_activation_chunk_fp32_to_fp16(__fp16 *restrict vtcm_dst, const float *restrict src, uint32_t n_rows, uint32_t k_block, uint32_t k_stride, uint32_t k_valid) {
const uint32_t n_rows_padded = hex_align_up(n_rows, HTP_MM_HMX_TILE_N_ROWS);
const uint32_t n_rows_tiled = (n_rows / HTP_MM_HMX_TILE_N_ROWS) * HTP_MM_HMX_TILE_N_ROWS;
uint32_t r = 0;
#pragma unroll(2)
for (r = 0; r < n_rows_tiled; r += 2) {
uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
const float *ptr_in0 = src + (r + 0) * k_stride;
const float *ptr_in1 = src + (r + 1) * k_stride;
uint32_t c = 0;
for (; c + 32 <= k_valid; c += 32) {
HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
tile[r1 / 2] = v_out;
}
if (c < k_block) {
HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
uint32_t rem = k_valid - c;
HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
tile[r1 / 2] = v_out;
}
}
for (; r < n_rows_padded; r += 2) {
uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
const bool row0_valid = r < n_rows;
const bool row1_valid = (r + 1) < n_rows;
const float *ptr_in0 = row0_valid ? (src + (r + 0) * k_stride) : NULL;
const float *ptr_in1 = row1_valid ? (src + (r + 1) * k_stride) : NULL;
uint32_t c = 0;
for (; c + 32 <= k_valid; c += 32) {
HVX_Vector v0 = Q6_V_vzero();
HVX_Vector v1 = Q6_V_vzero();
if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
tile[r1 / 2] = v_out;
}
if (c < k_block) {
HVX_Vector v0 = Q6_V_vzero();
HVX_Vector v1 = Q6_V_vzero();
if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
uint32_t rem = k_valid - c;
HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
tile[r1 / 2] = v_out;
}
}
}
static void transfer_activation_row_pair_fp32_to_fp16(
__fp16 *restrict vtcm_dst,
const float *restrict row0,
+2
View File
@@ -150,7 +150,9 @@ int op_matmul_nx(struct htp_ops_context * octx);
int op_matmul_id_nx(struct htp_ops_context * octx);
int op_binary(struct htp_ops_context * octx);
int op_unary(struct htp_ops_context * octx);
int op_sum(struct htp_ops_context * octx);
int op_sum_rows(struct htp_ops_context * octx);
int op_argmax(struct htp_ops_context * octx);
int op_activations(struct htp_ops_context * octx);
int op_softmax(struct htp_ops_context * octx);
int op_add_id(struct htp_ops_context * octx);
+4
View File
@@ -23,6 +23,7 @@ enum htp_data_type {
HTP_TYPE_Q4_1 = 3,
HTP_TYPE_Q8_0 = 8,
HTP_TYPE_Q4_K = 12,
HTP_TYPE_Q5_K = 13,
HTP_TYPE_Q6_K = 14,
HTP_TYPE_IQ4_NL = 20,
HTP_TYPE_I32 = 26,
@@ -68,6 +69,7 @@ enum htp_op_code {
HTP_OP_UNARY_ABS,
HTP_OP_UNARY_LOG,
HTP_OP_UNARY_RELU,
HTP_OP_UNARY_STEP,
HTP_OP_GLU_SWIGLU,
HTP_OP_GLU_SWIGLU_OAI,
HTP_OP_GLU_GEGLU,
@@ -85,6 +87,7 @@ enum htp_op_code {
HTP_OP_TOP_K,
HTP_OP_SQR,
HTP_OP_SQRT,
HTP_OP_SUM,
HTP_OP_SUM_ROWS,
HTP_OP_SSM_CONV,
HTP_OP_REPEAT,
@@ -107,6 +110,7 @@ enum htp_op_code {
HTP_OP_GLU_SWIGLU_CLAMP,
HTP_OP_MDEV_GROUP,
HTP_OP_ROLL,
HTP_OP_ARGMAX,
HTP_OP_INVALID
};
+52
View File
@@ -579,6 +579,58 @@ static inline void hvx_abs_f16(uint8_t * restrict dst, const uint8_t * restrict
}
}
//
// Step
//
static inline void hvx_step_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
assert((unsigned long) dst % 128 == 0);
assert((unsigned long) src % 128 == 0);
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
const uint32_t elem_size = sizeof(float);
const uint32_t epv = 128 / elem_size;
const uint32_t nvec = n / epv;
const uint32_t nloe = n % epv;
uint32_t i = 0;
_Pragma("unroll(4)")
for (; i < nvec; i++) {
vdst[i] = hvx_vec_step_f32(vsrc[i]);
}
if (nloe) {
HVX_Vector v = hvx_vec_step_f32(vsrc[i]);
hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v);
}
}
static inline void hvx_step_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
assert((unsigned long) dst % 128 == 0);
assert((unsigned long) src % 128 == 0);
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
const uint32_t elem_size = sizeof(_Float16);
const uint32_t epv = 128 / elem_size;
const uint32_t nvec = n / epv;
const uint32_t nloe = n % epv;
uint32_t i = 0;
_Pragma("unroll(4)")
for (; i < nvec; i++) {
vdst[i] = hvx_vec_step_f16(vsrc[i]);
}
if (nloe) {
HVX_Vector v = hvx_vec_step_f16(vsrc[i]);
hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v);
}
}
//
// Square
//
+14
View File
@@ -111,6 +111,20 @@ static inline HVX_Vector hvx_vec_neg_f32(HVX_Vector v) {
#endif // __HVX_ARCH__ > 75
}
static inline HVX_Vector hvx_vec_step_f32(HVX_Vector v) {
const HVX_Vector zero = Q6_V_vzero();
const HVX_Vector one = hvx_vec_splat_f32(1.0f);
HVX_VectorPred q = Q6_Q_vcmp_gt_VsfVsf(v, zero);
return Q6_V_vmux_QVV(q, one, zero);
}
static inline HVX_Vector hvx_vec_step_f16(HVX_Vector v) {
const HVX_Vector zero = Q6_V_vzero();
const HVX_Vector one = hvx_vec_splat_f16((_Float16) 1.0f);
HVX_VectorPred q = Q6_Q_vcmp_gt_VhfVhf(v, zero);
return Q6_V_vmux_QVV(q, one, zero);
}
static inline HVX_VectorPred hvx_vec_is_nan_f16(HVX_Vector v) {
const HVX_Vector vnan_exp = Q6_Vh_vsplat_R(0x7C00);
const HVX_Vector vnan_frac = Q6_Vh_vsplat_R(0x7FFF);
+212 -47
View File
@@ -121,29 +121,43 @@ static inline void quantize_block_f32_q8_0_tiled(float * restrict x, uint8_t * r
HVX_Vector * vx = (HVX_Vector *) x;
HVX_Vector zero = Q6_V_vzero();
HVX_Vector vmax0_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[0]));
HVX_Vector vmax1_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[1]));
HVX_Vector vmax2_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[2]));
HVX_Vector vmax3_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[3]));
HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero);
HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero);
HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero);
HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero);
HVX_Vector vmax0_qf = Q6_Vqf32_vsub_VsfVsf(vmax0_sf, zero);
HVX_Vector vmax1_qf = Q6_Vqf32_vsub_VsfVsf(vmax1_sf, zero);
HVX_Vector vmax2_qf = Q6_Vqf32_vsub_VsfVsf(vmax2_sf, zero);
HVX_Vector vmax3_qf = Q6_Vqf32_vsub_VsfVsf(vmax3_sf, zero);
HVX_Vector vmax01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax1_qf, vmax0_qf)));
HVX_Vector vmax23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax3_qf, vmax2_qf)));
HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf)));
HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf)));
HVX_Vector vmax_hf = hvx_vec_reduce_max_f16(hvx_vec_abs_f16(vx01_hf));
vmax_hf = hvx_vec_reduce_max2_f16(hvx_vec_abs_f16(vx23_hf), vmax_hf);
HVX_Vector vd01_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax01_hf, Q6_Vh_vsplat_R(0x2008));
HVX_Vector vd23_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax23_hf, Q6_Vh_vsplat_R(0x2008));
HVX_Vector vd01_hf = Q6_Vhf_equals_Vqf16(vd01_qf16);
HVX_Vector vd23_hf = Q6_Vhf_equals_Vqf16(vd23_qf16);
HVX_Vector vd_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax_hf, Q6_Vh_vsplat_R(0x2008));
HVX_Vector vd_hf = Q6_Vhf_equals_Vqf16(vd_qf16);
HVX_Vector vd_inv_hf = hvx_vec_inverse_f16(vd_hf);
vx01_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd_inv_hf));
vx23_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd_inv_hf));
HVX_Vector vd01_inv_hf = hvx_vec_inverse_f16(vd01_hf);
HVX_Vector vd23_inv_hf = hvx_vec_inverse_f16(vd23_hf);
vx01_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd01_inv_hf));
vx23_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd23_inv_hf));
HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf);
HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf);
HVX_Vector vx_i8 = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16);
HVX_Vector r_scale = hvx_vec_repl_f16(vd_hf);
HVX_VectorPair vp01 = Q6_W_vshuff_VVR(vd01_hf, vd01_hf, -64);
HVX_VectorPair vp23 = Q6_W_vshuff_VVR(vd23_hf, vd23_hf, -64);
static const uint8_t __attribute__((aligned(128))) repl[128] = {
0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
@@ -157,8 +171,19 @@ static inline void quantize_block_f32_q8_0_tiled(float * restrict x, uint8_t * r
};
HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
#pragma unroll
for (int b = 0; b < 4; b++) {
HVX_Vector v_act = Q6_V_vror_VR(vx_i8, b * 32);
HVX_Vector r_scale;
if (b == 0) {
r_scale = Q6_V_lo_W(vp01);
} else if (b == 1) {
r_scale = Q6_V_hi_W(vp01);
} else if (b == 2) {
r_scale = Q6_V_lo_W(vp23);
} else {
r_scale = Q6_V_hi_W(vp23);
}
HVX_Vector r0 = Q6_V_vdelta_VV(v_act, v_repl_ctrl);
HVX_Vector r1 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 4), v_repl_ctrl);
@@ -371,6 +396,78 @@ static inline HVX_VectorPair accum_q8_0_32x2(
return Q6_W_vcombine_VV(v_sum1, v_sum0);
}
// Q5_K: OR 0x10 into every lane of v whose flag j is set in the plane (see HTP_MM_WEIGHT_TILE_SIZE_Q5_K)
static inline HVX_Vector hvx_q5k_or_hibit(HVX_Vector v, HVX_Vector v_plane, int j) {
HVX_VectorPred q = Q6_Q_vand_VR(v_plane, 0x01010101u << j);
return Q6_V_vandor_VQR(v, q, 0x10101010);
}
// 5-bit variant: the high bit comes from the plane, see hvx_q5k_or_hibit
static inline HVX_VectorPair unpack_and_interleave_5bit_x2(HVX_Vector v_src, HVX_Vector v_plane, int i, HVX_Vector mask_h4) {
HVX_Vector v_lo = hvx_q5k_or_hibit(Q6_V_vand_VV(v_src, mask_h4), v_plane, 2 * i);
HVX_Vector v_hi = hvx_q5k_or_hibit(Q6_Vub_vlsr_VubR(v_src, 4), v_plane, 2 * i + 1);
HVX_VectorPair v01_pair = Q6_W_vshuff_VVR(v_hi, v_lo, -1);
HVX_Vector v01_lo = Q6_V_lo_W(v01_pair);
HVX_Vector v01_hi = Q6_V_hi_W(v01_pair);
HVX_Vector v23_lo = Q6_V_valign_VVR(v01_hi, v01_lo, 64);
HVX_Vector v_W0 = Q6_V_lo_W(Q6_W_vshuff_VVR(v23_lo, v01_lo, -2));
HVX_Vector v67_lo = Q6_V_valign_VVR(v01_lo, v01_hi, 64);
HVX_Vector v_W1 = Q6_V_lo_W(Q6_W_vshuff_VVR(v67_lo, v01_hi, -2));
return Q6_W_vcombine_VV(v_W1, v_W0);
}
static inline HVX_Vector accum_5bit_32x1(
const HVX_Vector * restrict vptr,
const HVX_Vector * restrict v_act,
HVX_Vector i8
) {
HVX_Vector v_sum0 = Q6_V_vzero();
HVX_Vector v_sum1 = Q6_V_vzero();
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
HVX_Vector v_plane = vptr[5];
#pragma unroll
for (int i = 0; i < 4; i++) {
HVX_VectorPair v_W_pair = unpack_and_interleave_5bit_x2(vptr[i], v_plane, i, mask_h4);
HVX_Vector v_W0 = Q6_Vb_vsub_VbVb(Q6_V_lo_W(v_W_pair), i8);
HVX_Vector v_W1 = Q6_Vb_vsub_VbVb(Q6_V_hi_W(v_W_pair), i8);
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act[i * 2 + 0]);
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act[i * 2 + 1]);
}
return Q6_Vw_vadd_VwVw(v_sum0, v_sum1);
}
static inline HVX_VectorPair accum_5bit_32x2(
const HVX_Vector * restrict vptr,
const HVX_Vector * restrict v_act0,
const HVX_Vector * restrict v_act1,
HVX_Vector i8
) {
HVX_Vector v_sum0 = Q6_V_vzero();
HVX_Vector v_sum1 = Q6_V_vzero();
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
HVX_Vector v_plane = vptr[5];
#pragma unroll
for (int i = 0; i < 4; i++) {
HVX_VectorPair v_W_pair = unpack_and_interleave_5bit_x2(vptr[i], v_plane, i, mask_h4);
HVX_Vector v_W0 = Q6_Vb_vsub_VbVb(Q6_V_lo_W(v_W_pair), i8);
HVX_Vector v_W1 = Q6_Vb_vsub_VbVb(Q6_V_hi_W(v_W_pair), i8);
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act0[i * 2 + 0]);
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W1, v_act0[i * 2 + 1]);
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W0, v_act1[i * 2 + 0]);
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act1[i * 2 + 1]);
}
return Q6_W_vcombine_VV(v_sum1, v_sum0);
}
// Q6_K weights are stored unsigned (0..63), see HTP_MM_WEIGHT_TILE_SIZE_Q6_K. Unpack k-group g of a tile to signed bytes (q - 32)
static inline HVX_Vector unpack_q6_k_group(const HVX_Vector * restrict vptr, int g, HVX_Vector mask_0f, HVX_Vector mask_03, HVX_Vector i32) {
HVX_Vector v_lo = (g & 1) ? Q6_Vub_vlsr_VubR(vptr[g >> 1], 4) : Q6_V_vand_VV(vptr[g >> 1], mask_0f);
@@ -689,6 +786,102 @@ static void tiled_vec_dot_q8_0_32x2(const uint32_t n, float * restrict s0, float
}
}
static void tiled_vec_dot_q5_k_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
const uint8_t * restrict tile_ptr = vx;
const uint8_t * restrict y_q = vy;
HVX_Vector v_sum_float = Q6_V_vzero();
uint32_t n_k_tiles = n / 32;
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 768);
const HVX_Vector * restrict v_act = (const HVX_Vector *) (y_q + kt * 1280);
HVX_Vector v_sum = accum_5bit_32x1(vptr, v_act, Q6_V_vzero());
HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
HVX_Vector v_scale_offset = vptr[4];
HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
HVX_Vector v_scale = Q6_V_lo_W(p_deal);
HVX_Vector v_offset = Q6_V_hi_W(p_deal);
HVX_Vector v_scale_a = v_act[8];
HVX_Vector v_sum_a = v_act[9];
HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a);
HVX_Vector v_offset_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a);
HVX_Vector v_scaled_dot = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
HVX_Vector v_sum_scaled = hvx_vec_add_f32_f32(v_scaled_dot, v_offset_comb);
v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
}
if (sz) {
hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
} else {
hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
}
}
static void tiled_vec_dot_q5_k_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
const uint8_t * restrict tile_ptr = vx;
const uint8_t * restrict y0_q = vy0;
const uint8_t * restrict y1_q = vy1;
HVX_Vector v_sum_float_c0 = Q6_V_vzero();
HVX_Vector v_sum_float_c1 = Q6_V_vzero();
uint32_t n_k_tiles = n / 32;
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 768);
const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1280);
const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1280);
HVX_VectorPair v_sums = accum_5bit_32x2(vptr, v_act0, v_act1, Q6_V_vzero());
HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
HVX_Vector v_scale_offset = vptr[4];
HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
HVX_Vector v_scale = Q6_V_lo_W(p_deal);
HVX_Vector v_offset = Q6_V_hi_W(p_deal);
HVX_Vector v_scale_a_c0 = v_act0[8];
HVX_Vector v_sum_a_c0 = v_act0[9];
HVX_Vector v_scale_a_c1 = v_act1[8];
HVX_Vector v_sum_a_c1 = v_act1[9];
HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a_c0);
HVX_Vector v_offset_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a_c0);
HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a_c1);
HVX_Vector v_offset_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a_c1);
HVX_Vector v_scaled_dot_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
HVX_Vector v_sum_scaled_c0 = hvx_vec_add_f32_f32(v_scaled_dot_c0, v_offset_comb_c0);
HVX_Vector v_scaled_dot_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
HVX_Vector v_sum_scaled_c1 = hvx_vec_add_f32_f32(v_scaled_dot_c1, v_offset_comb_c1);
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
}
if (sz0) {
hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
} else {
hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
}
if (sz1) {
hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
} else {
hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
}
}
static void tiled_vec_dot_q6_k_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
const uint8_t * restrict tile_ptr = vx;
const uint8_t * restrict y_q = vy;
@@ -941,14 +1134,9 @@ static inline void quantize_f32_q8_0_tiled_kernel(
size_t src_row_size,
size_t dst_row_size
) {
const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
(void) tmp_data;
for (uint32_t i = 0; i < nrows; ++i) {
hex_l2fetch(src_data, src_row_size, src_row_size, 2);
hvx_copy_f32_aa(tmp_data, src_data, ne0);
quantize_row_f32_q8_0_tiled((float *) tmp_data, dst_data, ne0);
quantize_row_f32_q8_0_tiled((float *) src_data, dst_data, ne0);
dst_data += dst_row_size;
src_data += src_row_size;
}
@@ -963,14 +1151,9 @@ static inline void quantize_f32_q8_1_tiled_kernel(
size_t src_row_size,
size_t dst_row_size
) {
const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
(void) tmp_data;
for (uint32_t i = 0; i < nrows; ++i) {
hex_l2fetch(src_data, src_row_size, src_row_size, 2);
hvx_copy_f32_aa(tmp_data, src_data, ne0);
quantize_row_f32_q8_1_tiled((float *) tmp_data, dst_data, ne0);
quantize_row_f32_q8_1_tiled((float *) src_data, dst_data, ne0);
dst_data += dst_row_size;
src_data += src_row_size;
}
@@ -988,24 +1171,15 @@ static inline void quantize_f32_q8_0_tiled_block_kernel(
uint32_t r,
uint32_t c
) {
(void) tmp_data;
const uint32_t qk = QK_Q8_0_TILED;
const uint32_t nb = (ne0 + qk - 1) / qk;
for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
const float * restrict src_ptr = (const float *) ((const uint8_t *) src + r * src_row_size + c * qk * sizeof(float));
uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1152;
hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
if (c == nb - 1) {
uint32_t active_elements = ne0 - c * qk;
hvx_splat_f32_a(tmp_data, 0.0f, qk);
hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
} else {
hvx_copy_f32_aa(tmp_data, src_ptr, qk);
}
quantize_block_f32_q8_0_tiled((float *) tmp_data, dst_ptr);
quantize_block_f32_q8_0_tiled((float *) src_ptr, dst_ptr);
c++;
if (c == nb) {
@@ -1027,24 +1201,15 @@ static inline void quantize_f32_q8_1_tiled_block_kernel(
uint32_t r,
uint32_t c
) {
(void) tmp_data;
const uint32_t qk = QK_Q8_0_TILED;
const uint32_t nb = (ne0 + qk - 1) / qk;
for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
const float * restrict src_ptr = (const float *) ((const uint8_t *) src + r * src_row_size + c * qk * sizeof(float));
uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1280;
hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
if (c == nb - 1) {
uint32_t active_elements = ne0 - c * qk;
hvx_splat_f32_a(tmp_data, 0.0f, qk);
hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
} else {
hvx_copy_f32_aa(tmp_data, src_ptr, qk);
}
quantize_block_f32_q8_1_tiled((float *) tmp_data, dst_ptr);
quantize_block_f32_q8_1_tiled((float *) src_ptr, dst_ptr);
c++;
if (c == nb) {
+73
View File
@@ -326,6 +326,79 @@ static inline int32_t hvx_reduce_max_i32(const uint8_t * restrict src, const int
}
}
static inline void hvx_argmax_f32(
const float * restrict src,
uint32_t n,
uint32_t offset,
float * out_val,
int32_t * out_idx
) {
if (n == 0) {
*out_val = -INFINITY;
*out_idx = (int32_t) offset;
return;
}
if (n < 32 || !hex_is_aligned((void *) src, 128)) {
float best_val = src[0];
int32_t best_idx = (int32_t) offset;
for (uint32_t i = 1; i < n; i++) {
if (src[i] > best_val) {
best_val = src[i];
best_idx = (int32_t) (offset + i);
}
}
*out_val = best_val;
*out_idx = best_idx;
return;
}
static const int32_t c_lane_idx[32] __attribute__((aligned(128))) = {
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15,
16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31
};
const HVX_Vector v_init_idx = *(const HVX_Vector *) c_lane_idx;
const HVX_Vector v_step = Q6_V_vsplat_R(32);
HVX_Vector v_cur_idx = Q6_Vw_vadd_VwVw(v_init_idx, Q6_V_vsplat_R((int32_t) offset));
HVX_Vector v_max_val = hvx_vec_splat_f32(-INFINITY);
HVX_Vector v_max_idx = v_cur_idx;
const HVX_Vector * vsrc = (const HVX_Vector *) src;
const uint32_t nvec = n / 32;
for (uint32_t vi = 0; vi < nvec; vi++) {
HVX_Vector v = vsrc[vi];
HVX_VectorPred pred = Q6_Q_vcmp_gt_VsfVsf(v, v_max_val);
v_max_val = Q6_V_vmux_QVV(pred, v, v_max_val);
v_max_idx = Q6_V_vmux_QVV(pred, v_cur_idx, v_max_idx);
v_cur_idx = Q6_Vw_vadd_VwVw(v_cur_idx, v_step);
}
HVX_VectorAlias u_val, u_idx;
u_val.v = v_max_val;
u_idx.v = v_max_idx;
float best_val = u_val.fp32[0];
int32_t best_idx = (int32_t) u_idx.w[0];
for (int i = 1; i < 32; i++) {
if (u_val.fp32[i] > best_val) {
best_val = u_val.fp32[i];
best_idx = (int32_t) u_idx.w[i];
}
}
for (uint32_t i = nvec * 32; i < n; i++) {
if (src[i] > best_val) {
best_val = src[i];
best_idx = (int32_t) (offset + i);
}
}
*out_val = best_val;
*out_idx = best_idx;
}
#undef hvx_reduce_loop_body
#undef HVX_REDUCE_MAX_OP
#undef HVX_REDUCE_SUM_OP
+7
View File
@@ -857,6 +857,7 @@ static int execute_op(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS:
case HTP_OP_UNARY_LOG:
case HTP_OP_UNARY_RELU:
case HTP_OP_UNARY_STEP:
case HTP_OP_L2_NORM:
return op_unary(octx);
@@ -882,6 +883,9 @@ static int execute_op(struct htp_ops_context * octx) {
case HTP_OP_GET_ROWS:
return op_get_rows(octx);
case HTP_OP_SUM:
return op_sum(octx);
case HTP_OP_SUM_ROWS:
return op_sum_rows(octx);
@@ -898,6 +902,9 @@ static int execute_op(struct htp_ops_context * octx) {
case HTP_OP_TOP_K:
return op_top_k(octx);
case HTP_OP_ARGMAX:
return op_argmax(octx);
case HTP_OP_SSM_CONV:
return op_ssm_conv(octx);
File diff suppressed because it is too large Load Diff
+88 -53
View File
@@ -25,6 +25,9 @@ extern "C" {
#define HTP_MM_WEIGHT_TILE_SIZE_Q8_0 1088
#define HTP_MM_WEIGHT_TILE_SIZE_IQ4_NL 576
#define HTP_MM_WEIGHT_TILE_SIZE_MXFP4 544
// Q5_K: the Q4_1 tile (640) followed by a 128-byte plane with the 5th bit of every quant, transposed so that
// plane byte l holds the eight flags of lane l: bit 2i = low nibble of nibble vector i, bit 2i+1 = high nibble
#define HTP_MM_WEIGHT_TILE_SIZE_Q5_K 768
// Q6_K native 6-bit tile (32 rows x 32 k), vrmpy-ready: byte 4*row+b of a vector holds k = 4*group+b
// vectors 0..3: low nibbles, vector i holds group 2i (low nibble) and group 2i+1 (high nibble)
// vectors 4..5: high 2 bits, vector m holds groups 4m..4m+3 at bit offsets 0,2,4,6
@@ -37,6 +40,7 @@ extern "C" {
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q8_0 1152
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_IQ4_NL 640
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_MXFP4 640
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q5_K 768
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q6_K 896
// --- Activation Tiled Block Sizes (including padding) ---
@@ -199,6 +203,8 @@ static inline uint32_t htp_mm_get_weight_tile_size(int weight_type) {
return HTP_MM_WEIGHT_TILE_SIZE_Q4_1;
case HTP_TYPE_Q8_0:
return HTP_MM_WEIGHT_TILE_SIZE_Q8_0;
case HTP_TYPE_Q5_K:
return HTP_MM_WEIGHT_TILE_SIZE_Q5_K;
case HTP_TYPE_Q6_K:
return HTP_MM_WEIGHT_TILE_SIZE_Q6_K;
case HTP_TYPE_MXFP4:
@@ -218,6 +224,8 @@ static inline uint32_t htp_mm_get_weight_aligned_tile_size(int weight_type) {
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q4_1;
case HTP_TYPE_Q8_0:
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q8_0;
case HTP_TYPE_Q5_K:
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q5_K;
case HTP_TYPE_Q6_K:
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q6_K;
case HTP_TYPE_MXFP4:
@@ -227,6 +235,11 @@ static inline uint32_t htp_mm_get_weight_aligned_tile_size(int weight_type) {
}
}
// weight types whose tiles carry a per-block offset (x = d * q + m): the activations need block sums (q8_1)
static inline bool htp_mm_weight_has_offset(int weight_type) {
return weight_type == HTP_TYPE_Q4_1 || weight_type == HTP_TYPE_Q4_K || weight_type == HTP_TYPE_Q5_K;
}
// --- Activation/Row Size Helpers ---
static inline size_t htp_mm_q8_0_tiled_row_size(uint32_t ne) {
const uint32_t ne_padded = ((ne + 127) / 128) * 128;
@@ -248,6 +261,7 @@ static inline size_t htp_mm_get_tiled_row_stride(int weight_type, uint32_t k) {
case HTP_TYPE_Q4_1:
case HTP_TYPE_Q4_K:
case HTP_TYPE_Q8_0:
case HTP_TYPE_Q5_K:
case HTP_TYPE_Q6_K:
case HTP_TYPE_MXFP4:
return (size_t) nb * htp_mm_get_weight_tile_size(weight_type);
@@ -332,6 +346,7 @@ struct htp_mm_hvx_vtcm_layout {
size_t off_src2; // vtcm_src2 (Wq / fused only)
size_t off_src3; // vtcm_src3 (Wv / fused only)
size_t off_dst; // vtcm_dst (output scratch)
size_t off_act_raw; // vtcm_act_raw (raw activation DMA staging)
// Cached sizes
size_t src0_bytes;
@@ -339,6 +354,7 @@ struct htp_mm_hvx_vtcm_layout {
size_t src2_bytes;
size_t src3_bytes;
size_t dst_bytes;
size_t act_raw_bytes;
size_t total_bytes;
};
@@ -351,7 +367,6 @@ static inline void htp_mm_hmx_vtcm_layout_build(
size_t mc,
size_t nc,
uint32_t group_size,
bool use_dma_activation,
bool pipeline,
uint32_t act_threads,
uint32_t aligned_tile_size,
@@ -366,8 +381,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
const size_t activation_area_size = hex_align_up(group_size * act_head_stride * sizeof(uint16_t), HTP_MM_HMX_TILE_SIZE);
const size_t output_area_size = hex_align_up(group_size * mc * nc * sizeof(uint16_t), HTP_MM_HMX_TILE_SIZE);
const size_t scratch_area_size = hex_align_up(nc * vec_dot_size, HTP_MM_HMX_TILE_SIZE);
const size_t min_f32_size = use_dma_activation
? hex_align_up(act_threads * HTP_MM_DMA_ACT_MULTIPLIER * k * sizeof(float), 128) : 0;
const size_t min_f32_size = hex_align_up(act_threads * HTP_MM_DMA_ACT_MULTIPLIER * k * sizeof(float), 128);
// Group A: Permanent activation tiles and scales
size_t off_group_a = 0;
@@ -388,10 +402,9 @@ static inline void htp_mm_hmx_vtcm_layout_build(
// Group C: Activation prep temporary buffer (overlaps Group B, starting at off_group_a)
const size_t max_f32_size = act_threads * 64 * k * sizeof(float);
const size_t act_f32_size = use_dma_activation
? hex_align_up(hex_smin(max_f32_size, hex_smax(min_f32_size, group_b_size)), 128) : 0;
const size_t act_f32_size = hex_align_up(hex_smin(max_f32_size, hex_smax(min_f32_size, group_b_size)), 128);
size_t off_group_c = off_group_a;
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_c, off_act_f32, act_f32_size, use_dma_activation);
VTCM_LAYOUT_ALLOC(off_group_c, off_act_f32, act_f32_size);
const size_t group_c_size = off_group_c - off_group_a;
@@ -478,20 +491,20 @@ static inline void htp_mm_hvx_vtcm_layout_build(
bool is_fused_nx
) {
(void)src1_row_size;
size_t src0_sz = 0;
size_t src1_sz = 0;
size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0;
size_t src3_sz = 0;
size_t dst_sz = 0;
size_t src0_sz = 0;
size_t src1_sz = 0;
size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0;
size_t src3_sz = 0;
size_t dst_sz = 0;
size_t act_raw_sz = 0;
const bool is_repack = (wtype == HTP_TYPE_Q4_0 || wtype == HTP_TYPE_Q4_1 ||
wtype == HTP_TYPE_Q8_0 || wtype == HTP_TYPE_IQ4_NL ||
wtype == HTP_TYPE_MXFP4 || wtype == HTP_TYPE_Q6_K ||
wtype == HTP_TYPE_Q4_K);
wtype == HTP_TYPE_Q4_K || wtype == HTP_TYPE_Q5_K);
if (is_fused_nx) {
const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
const size_t quant_scratch_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads;
size_t weight_sz_per_thread = 0;
@@ -505,17 +518,19 @@ static inline void htp_mm_hvx_vtcm_layout_build(
weight_sz_per_thread = hex_round_up(n_prefetch * src0_row_size_padded, 128);
}
size_t tiled_act_row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
size_t tiled_act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
size_t act_sz = hex_round_up(tiled_act_row_size * src1_nrows, 128);
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
src1_sz = act_sz; // quantized activation buffer
src2_sz = 0;
src3_sz = 0;
dst_sz = quant_scratch_size;
src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
src1_sz = act_sz; // quantized activation buffer
src2_sz = 0;
src3_sz = 0;
dst_sz = 0;
act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
} else if (is_matmul_id) {
const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
const size_t src1_row_size_tiled = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10)
const size_t src1_row_size_tiled = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10)
: htp_mm_q8_0_tiled_row_size(ne10);
size_t src0_sz_per_thread = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
@@ -529,10 +544,13 @@ static inline void htp_mm_hvx_vtcm_layout_build(
src0_sz_per_thread = repacked_vtcm_size;
}
src0_sz = src0_sz_per_thread * n_threads;
dst_sz = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads;
src2_sz = 0;
src3_sz = 0;
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
src0_sz = src0_sz_per_thread * n_threads;
dst_sz = 0;
src2_sz = 0;
src3_sz = 0;
act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
} else {
const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
const size_t dst_nrows = (src1_nrows > 1) ? 0 : 1;
@@ -540,21 +558,23 @@ static inline void htp_mm_hvx_vtcm_layout_build(
switch (kernel_type) {
case HTP_MM_KERNEL_HVX_F16_F16_VTCM: {
size_t f16_src1_row_size = htp_mm_round_up(ne10 * 2, 128);
src1_sz = htp_mm_round_up(f16_src1_row_size * src1_nrows, 256);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
src1_sz = htp_mm_round_up(f16_src1_row_size * src1_nrows, 256);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
act_raw_sz = hex_round_up(hex_round_up(ne10 * sizeof(float), 128) * src1_nrows, 128);
break;
}
case HTP_MM_KERNEL_HVX_F32_F32_VTCM: {
size_t f32_src1_row_size = htp_mm_round_up(ne10 * 4, 128);
src1_sz = htp_mm_round_up(f32_src1_row_size * src1_nrows, 256);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
src1_sz = htp_mm_round_up(f32_src1_row_size * src1_nrows, 256);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
act_raw_sz = 0;
break;
}
case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
case HTP_MM_KERNEL_HVX_QUANT_ROW: {
size_t q_src1_row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
size_t q_src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
src1_sz = htp_mm_round_up(q_src1_row_size * src1_nrows, 256);
@@ -569,10 +589,10 @@ static inline void htp_mm_hvx_vtcm_layout_build(
src0_sz = repacked_vtcm_size * n_threads;
}
size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
size_t dst_slice_per_thread = (dst_nrows > 0 && src1_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread;
dst_sz = dst_size_per_thread * n_threads;
dst_sz = dst_slice_per_thread * n_threads;
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
break;
}
default:
@@ -580,19 +600,30 @@ static inline void htp_mm_hvx_vtcm_layout_build(
}
}
size_t off = 0;
VTCM_LAYOUT_ALLOC(off, off_src0, src0_sz);
VTCM_LAYOUT_ALLOC(off, off_src1, src1_sz);
VTCM_LAYOUT_ALLOC(off, off_src2, src2_sz);
VTCM_LAYOUT_ALLOC(off, off_src3, src3_sz);
VTCM_LAYOUT_ALLOC(off, off_dst, dst_sz);
// Group A: Persistent buffers across chunk compute
size_t off_group_a = 0;
VTCM_LAYOUT_ALLOC(off_group_a, off_src1, src1_sz);
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, src2_sz, src2_sz > 0);
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src3, src3_sz, src3_sz > 0);
L->src0_bytes = src0_sz;
L->src1_bytes = src1_sz;
L->src2_bytes = src2_sz;
L->src3_bytes = src3_sz;
L->dst_bytes = dst_sz;
L->total_bytes = off;
// Group B: Compute-only buffers (starts at off_group_a)
size_t off_group_b = off_group_a;
VTCM_LAYOUT_ALLOC(off_group_b, off_src0, src0_sz);
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_b, off_dst, dst_sz, dst_sz > 0);
const size_t group_b_size = off_group_b - off_group_a;
// Group C: Raw activation staging buffer (overlaps Group B, starts at off_group_a)
size_t off_group_c = off_group_a;
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_c, off_act_raw, act_raw_sz, act_raw_sz > 0);
const size_t group_c_size = off_group_c - off_group_a;
L->src0_bytes = src0_sz;
L->src1_bytes = src1_sz;
L->src2_bytes = src2_sz;
L->src3_bytes = src3_sz;
L->dst_bytes = dst_sz;
L->act_raw_bytes = act_raw_sz;
L->total_bytes = off_group_a + hex_smax(group_b_size, group_c_size);
}
static inline bool htp_mm_hvx_solve_vtcm_params(
@@ -630,7 +661,7 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
const size_t avail_act = vtcm_budget - fixed_bytes;
size_t row_size = 0;
if (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K)
row_size = htp_mm_weight_has_offset(wtype)
? htp_mm_q8_1_tiled_row_size(ne10)
: htp_mm_q8_0_tiled_row_size(ne10);
} else if (kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM) {
@@ -642,7 +673,12 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
return false;
}
uint32_t m_chunk = (uint32_t) (avail_act / row_size);
size_t eff_row_size = row_size;
if (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
eff_row_size += hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
}
uint32_t m_chunk = (uint32_t) (avail_act / eff_row_size);
if (m_chunk > 1) {
m_chunk &= ~1U;
}
@@ -679,15 +715,15 @@ static inline size_t htp_mm_hmx_get_2d_vtcm_size(
int wtype, uint32_t k, size_t mc, size_t nc, bool pipeline, uint32_t act_threads, uint32_t aligned_tile_size, size_t src2_size
) {
struct htp_mm_hmx_vtcm_layout L;
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, false, pipeline, act_threads, aligned_tile_size, src2_size);
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, pipeline, act_threads, aligned_tile_size, src2_size);
return L.total_bytes;
}
static inline size_t htp_mm_hmx_get_batched_vtcm_size(
int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool use_dma_activation, bool pipeline, uint32_t act_threads, size_t src2_size) {
int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool pipeline, uint32_t act_threads, size_t src2_size) {
(void)pipeline;
struct htp_mm_hmx_vtcm_layout L;
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, use_dma_activation, false, act_threads, 0, src2_size);
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, false, act_threads, 0, src2_size);
return L.total_bytes;
}
@@ -697,7 +733,6 @@ static inline bool htp_mm_hmx_solve_batched_params(
uint32_t ne01_padded,
uint32_t ne11,
uint32_t group_size,
bool use_dma_activation,
int n_threads,
bool pipeline,
size_t src2_size,
@@ -726,7 +761,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
if (htp_mm_hmx_compute_chunks(vtcm_budget, group_overhead, group_size_per_n, group_size_per_m, group_size_per_mn, hex_align_up(ne11, 32), ne01_padded,
(size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) ne11 * HTP_MM_HMX_COST_A_CONVERT,
&m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, use_dma_activation, pipeline, act_threads, src2_size);
size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, pipeline, act_threads, src2_size);
if (exact_size <= vtcm_budget) {
size_t mblocks = ((size_t) ne11 + m_chunk_candidate - 1) / m_chunk_candidate;
if (mblocks < best_mblocks || (mblocks == best_mblocks && act_threads > best_act_threads)) {
+260 -4
View File
@@ -83,16 +83,17 @@ static void sum_rows_thread_f32(unsigned int nth, unsigned int ith, void *data)
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_row);
for (uint32_t ir = 0; ir < n_rows; ir++) {
const float * restrict src_local = src_th + (ir * (src_stride / sizeof(float)));
const float * restrict src_local = (const float *) ((const uint8_t *) src_th + ir * src_stride);
float * restrict dst_local = (float *) ((uint8_t *) dst_th + ir * dst_stride);
if (ir + 1 < n_rows) {
hex_l2fetch(src_local + (src_stride / sizeof(float)), src_stride, src_stride, 1);
hex_l2fetch((const uint8_t *) src_local + src_stride, src_stride, src_stride, 1);
}
if (opt_path) {
dst_th[ir] = hvx_reduce_sum_f32_a((const uint8_t *) src_local, ne00);
*dst_local = hvx_reduce_sum_f32_a((const uint8_t *) src_local, ne00);
} else {
dst_th[ir] = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
*dst_local = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
}
}
@@ -152,3 +153,258 @@ int op_sum_rows(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
struct sum_context {
struct htp_ops_context * octx;
const float * src_data;
float partial_sums[HTP_MAX_NTHREADS];
uint32_t total_elems;
uint32_t elems_per_thread;
};
static void sum_thread_f32(unsigned int nth, unsigned int ith, void * data) {
struct sum_context * sctx = (struct sum_context *) data;
const uint32_t start = sctx->elems_per_thread * ith;
const uint32_t end = MIN(start + sctx->elems_per_thread, sctx->total_elems);
if (start >= end) {
sctx->partial_sums[ith] = 0.0f;
return;
}
const uint32_t n = end - start;
const float * src = sctx->src_data + start;
struct htp_thread_trace * tr = &sctx->octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
hex_l2fetch_block((const void *) src, n * sizeof(float));
sctx->partial_sums[ith] = hvx_reduce_sum_f32((const uint8_t *) src, n);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
}
int op_sum(struct htp_ops_context * octx) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * dst = octx->dst;
if (src0->type != HTP_TYPE_F32) {
return HTP_STATUS_NO_SUPPORT;
}
if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
return HTP_STATUS_NO_SUPPORT;
}
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
return HTP_STATUS_OK;
}
const uint32_t total_elems = (uint32_t) (src0->ne[0] * src0->ne[1] * src0->ne[2] * src0->ne[3]);
if (total_elems == 0) {
((float *) dst->data)[0] = 0.0f;
return HTP_STATUS_OK;
}
const uint32_t n_threads = (total_elems >= 1024) ? MIN(octx->n_threads, HTP_MAX_NTHREADS) : 1;
const uint32_t raw_chunk = (total_elems + n_threads - 1) / n_threads;
const uint32_t elems_per_thread = hex_round_up(raw_chunk, 32);
struct sum_context sctx = {
.octx = octx,
.src_data = (const float *) src0->data,
.total_elems = total_elems,
.elems_per_thread = elems_per_thread,
};
work_queue_run(octx->ctx->work_queue, sum_thread_f32, &sctx, n_threads);
float sum = 0.0f;
for (uint32_t i = 0; i < n_threads; i++) {
sum += sctx.partial_sums[i];
}
((float *) dst->data)[0] = sum;
return HTP_STATUS_OK;
}
static inline void argmax_slice_f32(
const float * restrict src,
uint32_t n,
uint32_t offset,
float * out_val,
int32_t * out_idx
) {
hvx_argmax_f32(src, n, offset, out_val, out_idx);
}
struct argmax_context {
struct htp_ops_context * octx;
const float * src_data;
int32_t * dst_data;
uint32_t ne00;
uint32_t src_stride;
uint32_t dst_stride;
uint32_t row_start;
uint32_t nrows;
uint32_t rows_per_thread;
uint32_t elems_per_thread;
float partial_max[HTP_MAX_NTHREADS];
int32_t partial_idx[HTP_MAX_NTHREADS];
};
static void argmax_thread_single_row(unsigned int nth, unsigned int ith, void * data) {
struct argmax_context * actx = (struct argmax_context *) data;
const uint32_t start = actx->elems_per_thread * ith;
const uint32_t end = MIN(start + actx->elems_per_thread, actx->ne00);
if (start >= end) {
actx->partial_max[ith] = -INFINITY;
actx->partial_idx[ith] = 0;
return;
}
const uint32_t n = end - start;
const float * src = actx->src_data + start;
struct htp_thread_trace * tr = &actx->octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
hex_l2fetch_block((const void *) src, n * sizeof(float));
float max_val;
int32_t max_idx;
argmax_slice_f32(src, n, start, &max_val, &max_idx);
actx->partial_max[ith] = max_val;
actx->partial_idx[ith] = max_idx;
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
}
static void argmax_thread_multi_row(unsigned int nth, unsigned int ith, void * data) {
struct argmax_context * actx = (struct argmax_context *) data;
const uint32_t r0 = actx->row_start + actx->rows_per_thread * ith;
const uint32_t r1 = MIN(r0 + actx->rows_per_thread, actx->row_start + actx->nrows);
if (r0 >= r1) {
return;
}
struct htp_thread_trace * tr = &actx->octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r0);
for (uint32_t r = r0; r < r1; r++) {
const float * src_row = (const float *) ((const uint8_t *) actx->src_data + r * actx->src_stride);
int32_t * dst_val = (int32_t *) ((uint8_t *) actx->dst_data + r * actx->dst_stride);
hex_l2fetch_block((const void *) src_row, actx->ne00 * sizeof(float));
float max_val;
int32_t max_idx;
argmax_slice_f32(src_row, actx->ne00, 0, &max_val, &max_idx);
*dst_val = max_idx;
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r0);
}
int op_argmax(struct htp_ops_context * octx) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * dst = octx->dst;
if (src0->type != HTP_TYPE_F32 || dst->type != HTP_TYPE_I32) {
return HTP_STATUS_NO_SUPPORT;
}
if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
return HTP_STATUS_NO_SUPPORT;
}
const uint32_t ne00 = src0->ne[0];
const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
if (ne00 == 0 || src0_nrows == 0) {
return HTP_STATUS_OK;
}
if (src0_nrows == 1) {
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
return HTP_STATUS_OK;
}
if (ne00 == 1) {
((int32_t *) dst->data)[0] = 0;
return HTP_STATUS_OK;
}
const uint32_t n_threads = (ne00 >= 1024) ? MIN(octx->n_threads, HTP_MAX_NTHREADS) : 1;
const uint32_t raw_chunk = (ne00 + n_threads - 1) / n_threads;
const uint32_t elems_per_thread = hex_round_up(raw_chunk, 32);
struct argmax_context actx = {
.octx = octx,
.src_data = (const float *) src0->data,
.dst_data = (int32_t *) dst->data,
.ne00 = ne00,
.src_stride = src0->nb[1] > 0 ? (uint32_t) src0->nb[1] : (uint32_t) (ne00 * sizeof(float)),
.dst_stride = dst->nb[0] > 0 ? (uint32_t) dst->nb[0] : (uint32_t) sizeof(int32_t),
.row_start = 0,
.nrows = 1,
.rows_per_thread = 1,
.elems_per_thread = elems_per_thread,
};
work_queue_run(octx->ctx->work_queue, argmax_thread_single_row, &actx, n_threads);
float best_val = actx.partial_max[0];
int32_t best_idx = actx.partial_idx[0];
for (uint32_t i = 1; i < n_threads; i++) {
if (actx.partial_max[i] > best_val) {
best_val = actx.partial_max[i];
best_idx = actx.partial_idx[i];
}
}
((int32_t *) dst->data)[0] = best_idx;
return HTP_STATUS_OK;
}
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) &&
htp_tensor_is_contiguous(dst, sizeof(int32_t));
const uint32_t elems_per_chunk = can_split ? 32 : 0;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
src0_nrows, elems_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
row_start = range.start;
nrows = range.count;
}
if (nrows == 0) {
return HTP_STATUS_OK;
}
const uint32_t n_threads = MIN(octx->n_threads, nrows);
const uint32_t rows_per_thread = (nrows + n_threads - 1) / n_threads;
struct argmax_context actx = {
.octx = octx,
.src_data = (const float *) src0->data,
.dst_data = (int32_t *) dst->data,
.ne00 = ne00,
.src_stride = src0->nb[1] > 0 ? (uint32_t) src0->nb[1] : (uint32_t) (ne00 * sizeof(float)),
.dst_stride = dst->nb[0] > 0 ? (uint32_t) dst->nb[0] : (uint32_t) sizeof(int32_t),
.row_start = row_start,
.nrows = nrows,
.rows_per_thread = rows_per_thread,
.elems_per_thread = 0,
};
work_queue_run(octx->ctx->work_queue, argmax_thread_multi_row, &actx, n_threads);
return HTP_STATUS_OK;
}
+38
View File
@@ -409,6 +409,20 @@ static void log_f16(const void * restrict src,
}
}
static void step_f16(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
const struct htp_unary_context * uctx) {
htp_unary_op_preamble;
for (uint32_t ir = 0; ir < num_rows; ir++) {
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
hvx_step_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0);
}
}
static void l2_norm_f16(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
@@ -662,6 +676,20 @@ static void relu_f32(const void * restrict src,
}
}
static void step_f32(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
const struct htp_unary_context * uctx) {
htp_unary_op_preamble;
for (uint32_t ir = 0; ir < num_rows; ir++) {
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
hvx_step_f32_aa(dst_local, src_local, ne0);
}
}
static void log_f32(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
@@ -767,6 +795,11 @@ static void tile_relu_f32(void * restrict dst, const void * restrict src, uint32
hvx_max_scalar_f32((uint8_t *) dst, (const uint8_t *) src, 0.0f, tw);
}
static void tile_step_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
(void) uctx;
hvx_step_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
}
static void tri_apply_tile_f32(const void * restrict src, void * restrict dst,
uint32_t tile_elems, uint32_t col_start, uint32_t i01,
uint32_t ne0, int32_t ttype) {
@@ -1486,6 +1519,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS: op_type = is_f16 ? "abs-f16" : "abs-f32"; break;
case HTP_OP_UNARY_LOG: op_type = is_f16 ? "log-f16" : "log-f32"; break;
case HTP_OP_UNARY_RELU: op_type = "relu-f32"; break;
case HTP_OP_UNARY_STEP: op_type = is_f16 ? "step-f16" : "step-f32"; break;
case HTP_OP_L2_NORM: op_type = is_f16 ? "l2norm-f16" : "l2norm-f32"; break;
case HTP_OP_TRI: op_type = "tri-f32"; break;
default:
@@ -1506,6 +1540,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_L2_NORM:
case HTP_OP_UNARY_ABS:
case HTP_OP_UNARY_LOG:
case HTP_OP_UNARY_STEP:
break;
default:
FARF(ERROR, "unary-%s: not supported for F16\n", op_type);
@@ -1634,6 +1669,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS: compute_func = (void *) tile_abs_f32; break;
case HTP_OP_UNARY_LOG: compute_func = (void *) tile_log_f32; break;
case HTP_OP_UNARY_RELU: compute_func = (void *) tile_relu_f32; break;
case HTP_OP_UNARY_STEP: compute_func = (void *) tile_step_f32; break;
case HTP_OP_TRI:
task_func = unary_thread_tiled_tri_f32;
compute_func = (void *) tri_apply_tile_f32;
@@ -1652,6 +1688,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_L2_NORM: compute_func = (void *) l2_norm_f16; break;
case HTP_OP_UNARY_ABS: compute_func = (void *) abs_f16; break;
case HTP_OP_UNARY_LOG: compute_func = (void *) log_f16; break;
case HTP_OP_UNARY_STEP: compute_func = (void *) step_f16; break;
default: break;
}
} else {
@@ -1678,6 +1715,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS: compute_func = (void *) abs_f32; break;
case HTP_OP_UNARY_LOG: compute_func = (void *) log_f32; break;
case HTP_OP_UNARY_RELU: compute_func = (void *) relu_f32; break;
case HTP_OP_UNARY_STEP: compute_func = (void *) step_f32; break;
case HTP_OP_L2_NORM: compute_func = (void *) l2_norm_f32; break;
case HTP_OP_TRI:
task_func = unary_thread_tri_f32;
+1
View File
@@ -59,6 +59,7 @@ static inline bool htp_op_is_unary(uint32_t opcode) {
case HTP_OP_UNARY_ABS:
case HTP_OP_UNARY_LOG:
case HTP_OP_UNARY_RELU:
case HTP_OP_UNARY_STEP:
case HTP_OP_L2_NORM:
case HTP_OP_TRI:
return true;

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