Compare commits

..
26 Commits
Author SHA1 Message Date
Adrien GallouëtandGitHub 3057bb66c8 ui : add cache (#28802)
Signed-off-by: Adrien Gallouët <angt@huggingface.co>
2026-09-12 16:09:46 +02:00
MiaoMing ChenandGitHub 56381e407c server : allow model downloads at model limit fix issue #26809 (#28530) 2026-09-12 11:50:35 +02:00
thelittlefiremanandGitHub c8edceb061 ggml-cuda: hip add specific config table for AMD GCN (#27841) 2026-09-12 11:26:53 +02:00
Adrien GallouëtandGitHub e192abb406 server : add missing headers (#28795)
Signed-off-by: Adrien Gallouët <angt@huggingface.co>
2026-09-12 11:23:54 +02:00
Alessandro de Oliveira Faria (A.K.A.CABELO)andGitHub 718f7b4175 vendor : update cpp-httplib to 0.56.0 (#28787) 2026-09-12 10:15:08 +03:00
Michael TaylorandGitHub 2a3005c23f syscl : Handle (fail gracefully) unsupported tq1_0 quants (#28681) 2026-09-12 03:05:38 -04:00
Ed AddarioandGitHub f3a33dff26 rpc : fix linking when compiling with BUILD_SHARED_LIBS=OFF (#28492) 2026-09-12 09:22:57 +03:00
PascalandGitHub c069aa7f5f server: frame the router child state command as a whole line (#28747)
The child writes its state commands on stdout while the logger writes
on stderr, and both share a single pipe. The logger emits the trailing
color reset after the newline of a debug, warn or error entry, so that
escape sequence has no newline of its own and the router reads it glued
in front of the next command. The line prefix check then fails and the
command is forwarded as a log line instead of being handled, which
leaves a finished download stuck in the downloading state.

Writing the command with a leading newline closes the pending line so
it always starts at a line boundary.
2026-09-12 07:38:50 +02:00
Hongqiang WangandGitHub 8a56aedd61 opencl: fix several bugs where the backend aborts (#27630) 2026-09-11 22:11:12 -07:00
07fc97716f opencl: add bin kernel kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8_bin (#28677)
* opencl: add A8 Q4_K non-MoE binary kernel

* opencl: fix layout compatibility

* opencl: rename binary kernel selection helpers

---------

Co-authored-by: Li He <lih@qti.qualcomm.com>
2026-09-11 22:10:08 -07:00
PascalandGitHub 3f5e94d7c2 webgpu: align tensor bindings to the type block size (#28382)
Walk the binding offset back until the distance to the tensor is a
whole number of blocks, so block quantized views get a valid element
offset in the shader.
2026-09-12 06:40:21 +02:00
eafe15a5e3 hexagon: support for multi-device model split (aka row-split) (#28589)
* hex-row-split: add support for multi-device row spliting

Co-authored-by: Max Krasnyansky <maxk@qti.qualcomm.com>

* hex-mdev: add work splitting to fused kernels

* hex-mdev: use mdev_ prefix for all multi-device state

* hex-mdev: make device configuration more expressive to support device groups

* hex-mdev: fix mdev session init

* hex-mdev: fused nx (2x,3x) matmuls must update row counts for each w/o

* hex-mdev: fix MUL_MAT work partitioning bugs introduced by mdev

* hex-cont: fix crashes with new tests due to wrong striding

* hex-mdev: move fences after l2flushes

* hex-cont: fix work splitting for mnpu -- align chunks to cachelines

* hex-mdev: fix CPY tests with multi-dev

* hex-mmid: fix work partitioning with mnpu

* hex-mm: fix test failures with mdev

* hex-binary: fix work partitioning for mdev

* hex-argsort: fix mdev partitioning

* hex-mdev: fix work partitioning and general updates for all simple ops

* hex-fa: fix mdev work splitting issues

* hex-mdev: fixing more failing ops test

* hex-mdev: update the rest of the ops

* hex-mdev: refactor all mdev splitting logic to be contained within if (mdev_count > 1) {...}

* hex-mdev: fix macros

* hex-mdev: simplify session flush logic

* hex-sync: fix recursion in session flush

* hex-mdev: factor out fence buffer and allocator

* hex-fence: make fence allocation more robust with reserved slots for mdev

* hex-mdev: keep all mdev state in htp_mdev_group

* hex-mdev: further cleanup mdev group handling at the host

* hex-mdev: update group idx in the opbatch before serializing

* hex-batch: remove separate op_pending and use batch_req/rsp_seq

* hex-async: workaround another missing tensor_init in ggml-meta

* hex-fence: cleanup and robustify fences and error handling in multi-device scenarios

* hex-ar: improve ALLREDUCE error handling

* hex-async: robust error handling for op_cpy_fence

* hex-async: use seq0 from allreduce context to allocate fence_seq

* hex-mdev: fix remaining issues with fence and barrier clearing in CPY_FENCE

* hex-misc: realign macros and fix misplaces trace events

* hex-misc: align macros

* hex-mdev: fix unclone buffer re-entrancy

* hex-glu: fix mdev partitioning logic

* hex-mdev: make buffer uncloning/cleanup work with tensor-split scenarios

* hex-mdev: tighten up the can_split check in act-ops

* hex-mdev: factor out common bits of the partitioning logic

* hex-mm: minor realignment of the macros

* hex-bufs: fix incorrectly placed assert for MAX_BUFS

* hex-pad: tighten up gating checks for PAD

* hex-kparams: make sure all kernels properly use kparams->n_threads

* hex-docs: update user and developer docs with new features and detailed guide for ops development

* hex-scripts: update run script to properly parse dev groups

* hex-misc: formatting

* hex-sess: minor cleanup for session init

* hex-ar: fix vtcm size calc in allreduce kparams

* hex-scripts: fix flake8 warnings

* hex-rope: update ROPE to support mdev work split

* hex-ops: remove redunant checks and minor reformat

* hex-dev-guide: update dev-guide to avoid redundant null checks

* hex-async: improve event_wait, event_sync and fence implementations

* hex-async: remove synchronous flush from event_sync

* hex-async: symplify fence recovery protocol and make sync more robust

* hex-async: futher simplify error recovery for fences

* hex-err: return status instead of just -1

* hex-async: print all seq nums in hex

* hex-async: make sure fences flush dirty ranges

* hex-async: add dirty ranges merging to reduce fence flushes

* hex-async: properly sync before freeing the event

* hex-async: make sure fence owner session is not overriden

* hex-async: more fence write order more robust

* hex-async: make sure not to fuse ALLREDUCE+ADD if their dsts overlap

* hex-fusion: cleanup redundant checks

---------

Co-authored-by: Alexander Lu <alexlu@qti.qualcomm.com>
2026-09-11 20:46:51 -07:00
d3146f2b56 ggml-webgpu: Update to a recent version of Dawn (#28683)
* ggml-webgpu: Update to a recent version of Dawn

* No module scanning

* Accept review suggestion to update comment

Co-authored-by: Masashi Yoshimura <yoshimura.masashi.frbs@gmail.com>

---------

Co-authored-by: Masashi Yoshimura <yoshimura.masashi.frbs@gmail.com>
2026-09-12 10:47:29 +09:00
82d6bb284d server: refactor subproc handling (#28555)
* server: refactor subproc handling

* fix Windows build

* download: keep concurrent downloads of one blob apart

Every process writes the same path + .downloadInProgress, so a second
download of the same blob finds that file, takes it for its own partial
transfer and asks for the bytes after it, which produces a corrupt
result. The in-progress file now carries the pid of the process writing
it.

std::rename also replaces an existing destination on POSIX but fails on
Windows, so a download whose blob appeared in the meantime is dropped
after every retry and an etag rewrite silently keeps the old value.
std::filesystem::rename has the POSIX behaviour everywhere, and the
error now carries the reason reported by the system.

* Revert "download: keep concurrent downloads of one blob apart"

This reverts commit 917b83f149.

* tests: serialize the router tests that download the same model

Parallel workers share one cache, so the two tests fetch the same blob
into the same in-progress file and race to rename it. They now take a
file lock around the download, like the session fixture does for the
preset models.

* Revert "tests: serialize the router tests that download the same model"

This reverts commit c368a4a98c.

---------

Co-authored-by: Pascal <admin@serveurperso.com>
2026-09-12 00:53:07 +02:00
8ea290247c cmake : skip PCH for llama-server PCH when using MSVC (#28763)
This commit fixes an issue that I introduced when adding PCH
(precompiled headers) in Commit 3bcfeb700f
("cmake : add PCH and unity build to improve build times (#28091)".

See linked issue for details.

Co-authored-by: mjungnickel18
Co-authored-by: Pascal <admin@serveurperso.com>

Resolves: https://github.com/ggml-org/llama.cpp/issues/28758
Refs: https://github.com/ggml-org/llama.cpp/actions/runs/34592933983/job/103262608990#step:9:1284
2026-09-11 21:36:52 +02:00
Georgi GerganovandGitHub b78a39a2f9 ci : run test-backend-ops as a dedicated ci/run.sh test (#28740)
* ci : run test-backend-ops as a dedicated gg test

Run test-backend-ops as a separate gg test in ci/run.sh so it is executed outside ctest. With GG_BUILD_HIGH_PERF it keeps the existing CPU-only invocation (-b CPU); otherwise it runs all available backends without a backend filter.

Remove the dedicated backend-ops workflow and keep test-backend-ops as a built target that is not registered with ctest to avoid duplicate runs.

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

* ci : run test-backend-ops earlier and enable high-perf on kleidiai

Move the test-backend-ops gg test before test-llama-archs.

Enable GG_BUILD_HIGH_PERF and LLAMA_ARG_THREADS on the Graviton4 KleidiAI job and use the standard self-hosted results/mnt paths.

Add TODO markers for decoupling tests from libllama.

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

* ci : run test-backend-ops in parallel

Pass -j $(nproc) to test-backend-ops in both high-perf and all-backend modes.

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

* ci : disable parallel tests for ROCm

* cont : disable parallel tests with MoltenVK
2026-09-11 22:00:57 +03:00
Rohanjames1997andGitHub 982937a333 tests: extend test-quantize-fns to test nrc=2 (i8mm) kernels (#16234)
* Test for nrc=2 as well | i8mm kernels

* Trigger only on supported HW

* Remove trailing whitespace

* Address review comment

* test: properly prepare nrc=2 inputs with independent data per row

* tests : make nrc=2 dot product inputs distinct

Assisted-by: Kiro

* tests : use non-trivial strides in nrc=2 dot product test

* tests : fail nrc=2 dot product test on non-finite errors
2026-09-12 02:19:37 +08:00
PascalandGitHub 8172e6577a tests: tolerate a shared pool abort in test_completion_unified (#28759)
The expected success table holds when the four requests enter the shared
pool together. On a loaded runner they are admitted tens of milliseconds
apart, the slot lifetimes overlap differently and the pool overflows
while a short request is still resident. The decode failure aborts every
slot, so a request the table marks as successful comes back with the
context error instead of its generation.

Such a request now passes on that error too, while any other status, a
different error or a truncated generation still fails the test.
2026-09-11 15:50:12 +02:00
Aman GuptaandGitHub 43f3dda623 ggml: skip 0-sized ids tensor when offloading selected experts (#28739) 2026-09-11 15:17:08 +02:00
Foad Abo DahoodandGitHub 5bda51bfbc metal : skip the empty half of the mul_mm_id token tile (#28301)
kernel_mul_mm_id splits its NR1 = 32 token tile into two 16-row halves and skips
the upper half when the expert did not fill it, on both the tensor and simdgroup
paths. The tB extents are corrected to (NK, NR1H) for the [NR1][NK] row-major tile.

The B tile is staged unconditionally, as on master: rows past nr1 restage a clamped
duplicate of a valid row, lie in the output-row dimension so they never contribute
to a valid row, and are dropped by the final store loop.

test-backend-ops: re-draw the expert ids between perf iterations of test_mul_mat_id
so MoE perf numbers are not warm-cache, and add token-tile boundary coverage using
n_used == n_mats, which routes every token to every expert so each expert receives
exactly n rows; n = 32, 33, 47, 48, 49 reach mul_mm_id and leave a last tile of 32,
1, 15, 16 and 17 rows.
2026-09-11 14:12:55 +03:00
Daniel BeveniusandGitHub 3bcfeb700f cmake : add PCH and unity build to improve build times (#28091)
* scripts : add initial profiling script (wip)

* src : add precompile headers (PCH) for models.h

* common : add common.h as PCH

* ggml : add PCH for ggml-impl.h

* mtmd : use PCH for models.h

* scripts : add script to build with Server/Tools/Tests

* server : add PCH for common.h

* docs: add profiling progress notes (wip)

* ggml : add exclude for GCC + SVE on ARM

Refs: https://github.com/ggml-org/llama.cpp/actions/runs/33393906061/job/99493756214?pr=28091

* ggml : attempt to fix use of std::hardware_destructive_inference_size

Refs: https://github.com/ggml-org/llama.cpp/actions/runs/33396221677/job/99501265689?pr=28091

* squash! ggml : attempt to fix use of std::hardware_destructive_inference_size

Add a version check for GCC 12 to conditionally apply the `-Winterference-size`
pragma.

* editorconfig : exclude profiling reports dir

This directory will not be included in the merge later and this commit
can be ignore at that point. Just fixing to keep CI happy.

* ggml : skip PCH for gcc on non-x86 architectures

* tests : add PCH for peg-parser/tests.h

There are 7 peg-parser tests that can share one PCH instead of then each
parsing the full tests.h.

* common : add PCH for chat.h

* docs : update linux build profiling full results

Just updating after a number of PCH additions. These are not exact
figures and will vary a bit from run to run, but they give a general idea
of the performance impact of PCH.

* cmake : introduce unity build for models

This commit introduces a unity build for the models to improve
compilation time.

The improvements were roughly the following:
```console
+------------------------+-----+------------+------------+------------+
| Build                  | TUs | Frontend   | Backend    | Total      |
+------------------------+-----+------------+------------+------------+
| Full,    master        | 396 |   811.0 s  |   692.2 s  | 1,503.2 s  |
| Full,    with PCH      | 405 |   380.0 s  |   664.7 s  | 1,044.7 s  |
| Full,    with PCH + UB | 264 |   357.7 s  |   635.7 s  |   993.4 s  |
+------------------------+-----+------------+------------+------------+

TU   = Translation Unit.
Full = includes Server, Tools, and Tests.
PCH  = precompiled headers.
UB   = unity build for models.
```

* docs : update linux profiling table with unitiy build results

* docs : update mac profiling results to include unity build [no ci]

* docs: remove profiling reports

* scripts : merge build profile scripts into one script

I was lazy before and just copied the first script to enable Tests,
Server, and Tools. This now merges them into a single script.

* Revert "editorconfig : exclude profiling reports dir" [no ci]

This reverts commit 2922a12118.

* src : rename ggml_view_2d_slice to gemma3n_view_2d_slice

This is to be consistent with the rename in gemma4.cpp which was
required to avoid a name clash.

* cmake : add build profile script for windows [no ci]

This commit adds a port of the scripts/build-profile.sh script to
windows powershell.

This was developed on Windows on ARM but should work on X64 as well but
needs to be tested there as well.
2026-09-11 13:01:29 +02:00
Daniel BeveniusandGitHub 1dfe94e048 common : fix typo in speculative.cpp comment [no ci] (#28750) 2026-09-11 12:59:43 +02:00
Georgi GerganovandGitHub a2878d30df metal : single-source fusion table + fusion debug rework (#28164)
* metal : rework fusion patterns into a single table

All fusable op patterns for the Metal backend are now declared once in a
fusion table (ggml-metal-fuse.cpp) and consumed by both the graph optimizer
(ggml_metal_fuse_max, packing) and the op encoders (ggml_metal_fuse_next,
compute). The two phases share the same pattern table plus ggml_can_fuse_subgraph_ext
for the structural checks, and differ only in the mode used for the pattern
check (STRUCTURAL at optimize time, since tensors are not allocated yet, and
FULL at compute time, including Metal buffer placement). This also protects the
snake activation (MUL + SIN + SQR + MUL + ADD) from being reordered during graph
optimization, which was previously unprotected.

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

* metal : fix absolute output indices in fusion patterns

ggml_can_fuse_subgraph_ext expects the outputs array to contain absolute graph
node indices (it indexes cgraph->nodes[outputs[i]]), but the fusion table query
was passing a relative index (n_ops - 1). As a result the last node of every
pattern was not recognized as an output and was subjected to the elidable
use-count check, which failed for essentially all fusions. This silently
disabled the norm/MUL fusion and caused a ~5% token-generation regression.

Pass the absolute graph index of the last node instead.

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

* metal : fuse gated_delta_net with cache cpy

Add GGML_METAL_FUSE_GDN_CACHE to the fusion table: when the gated_delta_net
kernel is followed by a cpy that scatters its recurrent state snapshots into
the KV cache, the kernel writes the snapshots straight into the cache buffer
and the trailing cpy is elided.

The gdn output has other consumers (the attn scores view), so unlike the
elision-chain patterns this is not a simple chain: a 'raw' flag on the fusion
pattern skips the generic chain/shape and ggml_can_fuse_subgraph_ext checks,
making the pattern-specific check callback the sole validator. Packing
(ggml_metal_fuse_max) now matches on the same view-transparent node sequence
that the compute phase uses, so the gdn + cache cpy group is packed along with
any intermediate views and stays adjacent through the reorder.

The fused cpy is a view consumer of the gdn (it writes the cache directly),
so its mem-range is skipped in the encoder; the skip is restricted to CPY
nodes consuming the previous fused node through a view so other fusions are
unaffected.

Add test_gated_delta_net_cache_fusion and register 5 cases.

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

* metal : drop is_view_consumer mem-range skip

The is_view_consumer skip was carried over from the upstream gated_delta_net
cache-fusion draft, but it is not needed: keeping the elided cpy's mem-range in
the concurrency tracker only ever adds a (conservative) memory barrier at the
fusion point. It can never remove a barrier, so it cannot introduce a race. The
worst case is one spurious barrier per gdn+cache-cpy fusion, which is within
run-to-run noise on Qwen3.5-0.8B Q8_0.

Dropping the check keeps the mem-range loop uniform for all fused groups.

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

* metal : rename gated_delta_net fused state output args

Rename the fused cache-write kernel argument to match the rest of the kargs:
state_out_stride -> nb_out (and widen it to uint64_t), and the local buffer id
bid_state_out -> bid_out.

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

* metal : rename raw fusion flag to unsafe

raw did not convey that the flag opts a fusion pattern out of the generic
elision-chain safety net (ggml_can_fuse_subgraph_ext + chain/shape checks).
rename it to 'unsafe' to make explicit that the pattern's check callback is the
sole validator and must re-establish the safety guarantees itself.

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

* metal : tidy fusion pattern checks and table

- const-correct ggml_metal_fuse_outputs buffer
- annotate unused check-callback parameters
- drop a redundant size_t cast
- align the ops/table initializers and add blank-line separation

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

* metal : add generic fusion stats via ad-hoc proc-address API

Add a device-owned fusion context that lets a test tool count how many
times each fusion pattern fires and toggle fusion. It is exposed through
the ad-hoc ggml_backend_reg_get_proc_address mechanism with generic names
so the testing tool is backend-agnostic:

- ggml_backend_fusion_stats_init: start collecting fusion stats; when a
  context is created afterwards it registers the labels/counters and
  encodes single-threaded (n_cb == 0) so the counters are race-free
- ggml_backend_fusion_stats_reset / _get_stats / _set_enabled

The context lives on the metal device (not on the last backend context),
so counters accumulate across contexts and reads are always consistent.
The enable/disable toggle is initialized from GGML_METAL_FUSION_DISABLE
and can be overridden by the test through set_enabled. Labels are
synthesized from the fuse table via ggml_metal_fuse_label (e.g.
"GATED_DELTA_NET+CPY").

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

* tests : add fusion count regression test with per-backend baseline

test-fusion runs every dummy model generated by test-llama-archs on a
single backend (single-threaded encoding, n_cb == 0) with fusion enabled
and disabled, and for each mode (prefill / decode) reports the per-fusion
counters and the NMSE between the fused and unfused logits, plus the NMSE
against a CPU reference.

A fusion pattern that silently stops matching (or fires when it should
not) is caught as a regression by comparing the counters against a
committed per-backend TSV baseline:

- --record writes the golden baseline, --check (default) validates it
- the unfused run doubles as a control: its counters must be all-zero
- NMSE is skipped when it is NaN or the arch is already broken on the
  device (e.g. plamo2 on Metal), so the count check is the hard gate
- baseline counts depend only on graph structure, not weights (verified
  stable across weight seeds)
- the fusion stats API is resolved through the ad-hoc get_proc_address
  mechanism with generic names; a backend that does not export it makes
  the test fail with an error

The committed MTL0.tsv baseline covers 110 dummy archs (298 rows).

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

* tests : rename fusion api helpers to match stats_init signature

Align the test with the ad-hoc fusion stats API: fusion_stats_init no
longer takes an enable bool (stats are turned on by calling it), so the
proc-address wrappers and typedefs are renamed to the api_* convention.

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

* tests : rename backend to device in fusion test CLI

The fusion test operates on a compute device (e.g. MTL0), not a backend,
so rename the --backend argument to --device and the backend_name
variable to device_name. Keep "backend" where it refers to the ggml
backend interface (the ad-hoc proc-address mechanism).

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

* tests : add --model and --help to fusion test

--model FILE runs the fusion regression test over a single model file
instead of enumerating a --models DIR. --models and --model are mutually
exclusive. Also add a --help/-h option that prints the usage.

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

* tests : use backend base name for fusion baseline output

The fusion test is invoked with a specific device name (e.g. MTL0), but
its output - the recorded baseline and the header it writes - should be
named after the backend base name (e.g. MTL, via ggml_backend_reg_name),
since the counters depend on the backend, not on the specific device
index. Rename the committed baseline to MTL.tsv.

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

* tests : run fusion test from ci instead of ctest

The fusion test needs Metal and generates a lot of dummy models, so it
does not belong in the generic ctest suite. Move it to ci/run.sh as
gg_run_test_fusion, gated on GG_BUILD_METAL like
gg_run_test_llama_archs_tensor_split: it generates the dummy models with
test-llama-archs -o and then validates the fusion counts against the
committed baseline. test-fusion.cpp is still built (llama_build) but no
longer registered as a ctest.

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

* tests : align fusion baseline TSV columns

Pad the TSV fields to fixed widths so the columns line up regardless of
the variable arch and fusion-label lengths, and trim each field on parse
so the padded file is still accepted. Regenerate the committed MTL.tsv
baseline in the padded format (data unchanged, verified identical modulo
padding).

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

* tests : widen label column and align fusion TSV header

Give the label column more room (28 chars) and fix the column header
widths so they match the data rows (moe/mode/label), keeping the header
aligned with the values. Regenerate the MTL.tsv baseline in the new
format (data unchanged).

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

* tests : switch fusion baseline from TSV to CSV

Use comma-separated values like the rest of the project, keeping the
padded, aligned columns. Split on ',' and trim on parse. Rename the
committed baseline to MTL.csv (data unchanged, verified identical modulo
padding/separator). Update the ci/run.sh check path accordingly.

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

* cont : rebase + update MTL stats

* tests : avoid graph reallocations for some archs

* metal : tidy fusion debugging context and op init

- simplify the shared fusion debugging context comments
- shorten the ggml_metal_fusion struct comment
- align the ggml_metal_fuse struct fields and comments
- move the fusion parameter of ggml_metal_op_init right after dev

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

* tests : dedup fusion baseline into any mode

prefill and decode always produce the same per-graph fusion count, so
store a single row per label with mode = "any" and the per-graph count
instead of two rows. this halves the baseline size and keeps the check
stable.

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

* ci : move fusion model generation to a separate step

the dummy models generated by test-llama-archs are reused by other tests,
so generate them once in their own step instead of inside test_fusion.

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

* tests : bump nmse thold

* models : fix plamo2 graph

* tests : remove "skip" logic from test-fusion

* tests : set qwen3tts dummy vocab to codec head size

the dummy qwen3tts model used a vocab of 4096 while the codec head is
3072, so the graph padded the output with -inf which made the NMSE in
test-fusion produce NaN. use the exact codec head size instead so the
padding is not generated at all.

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

* tests : regen fusion baseline

reflect the plamo2 graph fix, which changed its fusion pattern split
(RMS_NORM+MUL 11->10, RMS_NORM+MUL+ADD 3->4; same total).

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

* ci : skip dummy model generation on OpenVINO

test-llama-archs does not build on the OpenVINO platform, so do not try
to generate the dummy models there.

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

* cont : minor

* tests : enable test-llama-archs on windows

* cont : disable on windows + workaround

* metal : naming nits

* test-fusion : add instructions to update baseline

* context : fix Kimi-K3 graph reserve

* fusion : update MTL

* cont : fix naming

* metal : rework fusion info storage

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

* metal : align fusion info API

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

* metal : use opaque fusion handle in ad-hoc API

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

* ci : move fusion test to dedicated workflow

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

* cont : run only on ggml changes

* cont : simplify

* fusion : remove multi-output stuff for now

* ci : fix typo
2026-09-11 12:41:54 +03:00
Foad Abo DahoodandGitHub aac810230f metal : fix idle threads in the remaining iq mul_mv kernels for ne00 < 1024 (#28692)
* metal : fix idle threads in the remaining iq mul_mv kernels for ne00 < 1024

Generalize the row split from #28086 to the six other kernels that use the
same lane-to-block mapping: iq1_s, iq1_m, iq2_xxs, iq2_xs, iq2_s and iq3_s.

Each of them assigns one 32-element chunk per thread, so when a row has
fewer than 32 chunks the rest of the simdgroup is idle. When nb32 < 32 and
nb32 divides 32, 32/nb32 threads now share each chunk and each takes a
slice of the rows, reusing the FC_mul_mv_split function constant and the
dispatch wrapper introduced for iq3_xxs.

The plain path is untouched: wide matrices keep one thread per chunk and
N_R0_<TYPE> = 4. Only the split path uses N_R0_<TYPE>_SPLIT = 8. The
K-quants have the same idle-thread issue but a different lane mapping, so
they are left for a separate change.

* metal : offset the src0 row pointer once in the iq mul_mv kernels

q2, dh, sc, qh and signs are all derived from xr, so the row slice
offset only has to be applied to xr.

* metal : fold iq mul_mv row split into offset0

Compute row0 and row1 before initializing the source pointers and apply
the row slice directly to offset0.

This keeps x and its derived pointers on the existing path while applying
the split row offset once.
2026-09-11 12:30:20 +03:00
Logan ChuandGitHub 5cdd3d1dad model : fix MTP context kv cache allocation for deepseek2, glm4moe, … (#28630)
* model : fix MTP context kv cache allocation for deepseek2, glm4moe, cohere2moe architectures (#28626)

* model: add inverse architecture gating and comprehensive architecture testing for mtp layer filtering

* model : slim NextN filter comment, drop test-llama-archs changes
2026-09-11 12:02:31 +03:00
Jesus GulfoandGitHub b0dcb8192b server: fix speculation after an image (#28715)
* server: fix speculation after an image

Pass the actual position to the drafter after an image, instead of the
token count. Affects every drafter, not just DFlash.

* rename draft n_past to pos0

n_past is used to denote number of tokens and this parameter is meant to be a position
2026-09-11 11:33:26 +03:00
117 changed files with 8994 additions and 2956 deletions
+1 -1
View File
@@ -33,7 +33,7 @@ env:
LLAMA_ARG_LOG_PREFIX: 1
LLAMA_ARG_LOG_TIMESTAMPS: 1
# TODO: fix failing tests on OpenVINO backend
CTEST_EXCLUDE: "test-llama-archs|^test-recurrent-state-|test-backend-ops|test-save-load-state"
CTEST_EXCLUDE: "test-llama-archs|^test-recurrent-state-|test-save-load-state"
jobs:
ubuntu-24-openvino:
+12 -6
View File
@@ -160,10 +160,10 @@ jobs:
- name: Dawn Dependency
id: dawn-depends
run: |
DAWN_VERSION="v20260317.182325"
DAWN_VERSION="v20260908.214631"
DAWN_OWNER="google"
DAWN_REPO="dawn"
DAWN_ASSET_NAME="Dawn-18eb229ef5f707c1464cc581252e7603c73a3ef0-ubuntu-latest-Release"
DAWN_ASSET_NAME="Dawn-94c3c9cc0d5fb2e85aebb370fa8d37b71aa34655-ubuntu-latest-Release"
echo "Fetching release asset from https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
curl -L -o artifact.tar.gz \
"https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
@@ -246,10 +246,10 @@ jobs:
- name: Dawn Dependency
id: dawn-depends
run: |
DAWN_VERSION="v20260317.182325"
DAWN_VERSION="v20260908.214631"
DAWN_OWNER="google"
DAWN_REPO="dawn"
DAWN_ASSET_NAME="Dawn-18eb229ef5f707c1464cc581252e7603c73a3ef0-macos-latest-Release"
DAWN_ASSET_NAME="Dawn-94c3c9cc0d5fb2e85aebb370fa8d37b71aa34655-macos-latest-Release"
echo "Fetching release asset from https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
curl -L -o artifact.tar.gz \
"https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
@@ -395,7 +395,11 @@ jobs:
- name: Test
id: ggml-ci
run: |
LLAMA_ARG_THREADS=$(nproc) GG_BUILD_HIGH_PERF=1 GG_BUILD_NO_BF16=1 GG_BUILD_EXTRA_TESTS_0=1 bash ./ci/run.sh ~/results/llama.cpp ~/mnt/llama.cpp
LLAMA_ARG_THREADS=$(nproc) \
GG_BUILD_HIGH_PERF=1 \
GG_BUILD_NO_BF16=1 \
GG_BUILD_EXTRA_TESTS_0=1 \
bash ./ci/run.sh ~/results/llama.cpp ~/mnt/llama.cpp
cpu-arm64-graviton4-kleidiai:
runs-on: ah-ubuntu_22_04-c8g_8x
@@ -434,6 +438,8 @@ jobs:
- name: Test
id: ggml-ci
run: |
LLAMA_ARG_THREADS=$(nproc) \
GG_BUILD_KLEIDIAI=1 \
GG_BUILD_EXTRA_TESTS_0=1 \
bash ./ci/run.sh ./tmp/results ./tmp/mnt
GG_BUILD_HIGH_PERF=1 \
bash ./ci/run.sh ~/results/llama.cpp ~/mnt/llama.cpp
+1 -3
View File
@@ -164,9 +164,7 @@ jobs:
export GGML_VK_VISIBLE_DEVICES=0
export GGML_VK_DISABLE_F16=1
export GGML_VK_DISABLE_COOPMAT=1
# This is using llvmpipe and runs slower than other backends
# test-backend-ops is too slow on llvmpipe, skip it
ctest -L main -E test-backend-ops --verbose --timeout 900
ctest -L main --verbose --timeout 900
windows:
runs-on: windows-2025
+1 -1
View File
@@ -68,7 +68,7 @@ jobs:
- name: Fetch emdawnwebgpu
run: |
DAWN_TAG="v20260317.182325"
DAWN_TAG="v20260908.214631"
EMDAWN_PKG="emdawnwebgpu_pkg-${DAWN_TAG}.zip"
echo "Downloading ${EMDAWN_PKG}"
curl -L -o emdawn.zip \
+5 -7
View File
@@ -77,10 +77,10 @@ jobs:
- name: Dawn Dependency
id: dawn-depends
run: |
DAWN_VERSION="v20260317.182325"
DAWN_VERSION="v20260908.214631"
DAWN_OWNER="google"
DAWN_REPO="dawn"
DAWN_ASSET_NAME="Dawn-18eb229ef5f707c1464cc581252e7603c73a3ef0-macos-latest-Release"
DAWN_ASSET_NAME="Dawn-94c3c9cc0d5fb2e85aebb370fa8d37b71aa34655-macos-latest-Release"
echo "Fetching release asset from https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
curl -L -o artifact.tar.gz \
"https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
@@ -147,10 +147,10 @@ jobs:
id: dawn-depends
run: |
sudo apt-get install -y libxrandr-dev libxinerama-dev libxcursor-dev mesa-common-dev libx11-xcb-dev libxi-dev
DAWN_VERSION="v20260317.182325"
DAWN_VERSION="v20260908.214631"
DAWN_OWNER="google"
DAWN_REPO="dawn"
DAWN_ASSET_NAME="Dawn-18eb229ef5f707c1464cc581252e7603c73a3ef0-ubuntu-latest-Release"
DAWN_ASSET_NAME="Dawn-94c3c9cc0d5fb2e85aebb370fa8d37b71aa34655-ubuntu-latest-Release"
echo "Fetching release asset from https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
curl -L -o artifact.tar.gz \
"https://github.com/google/dawn/releases/download/${DAWN_VERSION}/${DAWN_ASSET_NAME}.tar.gz"
@@ -190,6 +190,4 @@ jobs:
id: cmake_test
run: |
cd build
# This is using llvmpipe and runs slower than other backends
# test-backend-ops is too slow on llvmpipe, skip it
ctest -L main -E test-backend-ops --verbose --timeout 900
ctest -L main --verbose --timeout 900
+67
View File
@@ -0,0 +1,67 @@
name: Fusion
on:
workflow_dispatch: # allows manual triggering
push:
branches:
- master
paths: [
'.github/workflows/fusion.yml',
'ggml/**',
'tests/fusion/**',
'tests/test-fusion.cpp'
]
pull_request:
types: [opened, synchronize, reopened]
paths: [
'.github/workflows/fusion.yml',
'ggml/**',
'tests/fusion/**',
'tests/test-fusion.cpp'
]
concurrency:
group: ${{ github.workflow }}-${{ github.head_ref && github.ref || github.run_id }}
cancel-in-progress: true
env:
GGML_NLOOP: 3
GGML_N_THREADS: 1
LLAMA_ARG_LOG_COLORS: 1
LLAMA_ARG_LOG_PREFIX: 1
LLAMA_ARG_LOG_TIMESTAMPS: 1
jobs:
# TODO: add jobs for other backends as they adopt the fusion debug API
metal:
runs-on: [self-hosted, macOS, ARM64]
steps:
- name: Clone
id: checkout
uses: actions/checkout@v6
- name: Build
id: cmake_build
run: |
cmake -B build \
-DCMAKE_BUILD_TYPE=Release \
-DLLAMA_FATAL_WARNINGS=ON \
-DLLAMA_OPENSSL=OFF \
-DGGML_SCHED_NO_REALLOC=ON \
-DGGML_BLAS=OFF \
-DGGML_METAL=ON
time cmake --build build --config Release --target test-llama-archs -j $(sysctl -n hw.logicalcpu)
time cmake --build build --config Release --target test-fusion -j $(sysctl -n hw.logicalcpu)
- name: Generate models
id: generate_models
run: |
rm -rf build-ci-models && mkdir -p build-ci-models
./build/bin/test-llama-archs -o build-ci-models
- name: Test fusion
id: test_fusion
run: |
./build/bin/test-fusion --models build-ci-models --device MTL0 --check tests/fusion/MTL.csv
+57 -11
View File
@@ -190,7 +190,7 @@ if [ ! -z ${GG_BUILD_OPENVINO} ]; then
CMAKE_EXTRA="${CMAKE_EXTRA} -DGGML_OPENVINO=ON"
# TODO: fix failing tests on OpenVINO backend
CTEST_EXTRA="-E test-llama-archs|^test-recurrent-state-|test-backend-ops|test-save-load-state"
CTEST_EXTRA="-E test-llama-archs|^test-recurrent-state-|test-save-load-state"
fi
## helpers
@@ -250,7 +250,7 @@ function gg_run_ctest_debug {
(cmake -G "${CMAKE_GENERATOR}" -DCMAKE_BUILD_TYPE=Debug ${CMAKE_EXTRA} .. ) 2>&1 | tee -a $OUT/${ci}-cmake.log
(time cmake --build . --config Debug -j$(nproc)) 2>&1 | tee -a $OUT/${ci}-make.log
(time ctest -C Debug --output-on-failure -L main -E "test-opt|test-backend-ops|test-llama-archs" ${CTEST_EXTRA}) 2>&1 | tee -a $OUT/${ci}-ctest.log
(time ctest -C Debug --output-on-failure -L main -E "test-opt|test-llama-archs" ${CTEST_EXTRA}) 2>&1 | tee -a $OUT/${ci}-ctest.log
set +e
}
@@ -334,6 +334,35 @@ function gg_sum_test_llama_archs_tensor_split {
gg_printf '```\n'
}
# test_llama_archs_models
function gg_run_test_llama_archs_models {
cd ${SRC}
set -e
# TODO: fix and re-enable `test-llama-archs` on OpenVINO
# TODO: the `test-llama-archs` currently does not build on Windows, so we check if the binary exists
if [ -z ${GG_BUILD_OPENVINO} ] && [ -f ./build-ci-release/bin/test-llama-archs ]; then
rm -rf build-ci-models && mkdir -p build-ci-models
# generate the dummy models used by the model-dependent tests
./build-ci-release/bin/test-llama-archs -o build-ci-models 2>&1
fi
set +e
}
function gg_sum_test_llama_archs_models {
gg_printf '### %s\n\n' "${ci}"
gg_printf 'Generates the dummy models used by the model-dependent tests\n'
gg_printf '- status: %s\n' "$(cat $OUT/${ci}.exit)"
gg_printf '```\n'
gg_printf '%s\n' "$(cat $OUT/${ci}.log)"
gg_printf '```\n'
}
# test_scripts
function gg_run_test_scripts {
@@ -739,25 +768,43 @@ function gg_check_build_requirements {
fi
}
function gg_run_test_backend_ops_cpu {
function gg_run_test_backend_ops {
cd ${SRC}
cd build-ci-release
set -e
(time ./bin/test-backend-ops -b CPU ) 2>&1 | tee -a $OUT/${ci}-test-backend-ops-cpu.log
local args_extra="-j $(nproc)"
# TODO: fix multi-threaded for ROCm
# https://github.com/ggml-org/llama.cpp/actions/runs/34576278519/job/103297889044?pr=28740#step:3:4865
if [ ! -z ${GG_BUILD_ROCM} ]; then
args_extra=""
fi
# TODO: MoltenVK bug?
# https://github.com/ggml-org/llama.cpp/actions/runs/34611260059/job/103302413736?pr=28740#step:3:5897
if [ ! -z "${GG_BUILD_VULKAN}" ] && [ "$(uname -s)" = "Darwin" ]; then
args_extra=""
fi
if [ ! -z ${GG_BUILD_HIGH_PERF} ]; then
(time ./bin/test-backend-ops ${args_extra} -b CPU) 2>&1 | tee -a $OUT/${ci}-test-backend-ops.log
else
(time ./bin/test-backend-ops ${args_extra} ) 2>&1 | tee -a $OUT/${ci}-test-backend-ops.log
fi
set +e
}
function gg_sum_test_backend_ops_cpu {
function gg_sum_test_backend_ops {
gg_printf '### %s\n\n' "${ci}"
gg_printf 'Runs test-backend-ops for CPU backend\n'
gg_printf 'Runs test-backend-ops\n'
gg_printf '- status: %s\n' "$(cat $OUT/${ci}.exit)"
gg_printf '```\n'
gg_printf '%s\n' "$(cat $OUT/${ci}-test-backend-ops-cpu.log)"
gg_printf '%s\n' "$(cat $OUT/${ci}-test-backend-ops.log)"
gg_printf '```\n'
gg_printf '\n'
}
@@ -790,11 +837,10 @@ ret=0
test $ret -eq 0 && gg_run ctest_debug
test $ret -eq 0 && gg_run ctest_release
test $ret -eq 0 && gg_run test_llama_archs_tensor_split
test $ret -eq 0 && gg_run test_backend_ops
if [ ! -z ${GG_BUILD_HIGH_PERF} ]; then
test $ret -eq 0 && gg_run test_backend_ops_cpu
fi
test $ret -eq 0 && gg_run test_llama_archs_models
test $ret -eq 0 && gg_run test_llama_archs_tensor_split
if [ -z ${GG_BUILD_LOW_PERF} ]; then
test $ret -eq 0 && gg_run embd_bge_small
+2
View File
@@ -134,6 +134,8 @@ set_target_properties(${TARGET} PROPERTIES
target_include_directories(${TARGET} PUBLIC .)
target_link_libraries (${TARGET} PUBLIC vendor::nlohmann vendor::sheredom)
target_compile_features (${TARGET} PUBLIC cxx_std_17)
target_precompile_headers (${TARGET} PRIVATE common.h)
target_precompile_headers (${TARGET} PRIVATE chat.h)
if (LLAMA_SUBPROCESS)
target_compile_definitions(${TARGET} PUBLIC LLAMA_SUBPROCESS)
+11 -11
View File
@@ -296,7 +296,7 @@ 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.n_past, { seq_id }, true);
common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
}
int ret = llama_decode(ctx_dft, batch);
@@ -355,7 +355,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
continue;
}
common_batch_add(batch, id, dp.n_past + i + 1, { seq_id }, true);
common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
}
if (batch.n_tokens == 0) {
@@ -1197,7 +1197,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
common_sampler_reset(smpls[seq_id].get());
const int32_t n = (int32_t) dp.n_past;
const int32_t n = (int32_t) dp.pos0;
const int32_t n_draft = params.n_max;
@@ -1493,7 +1493,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
const int32_t n_tokens = batch_in.n_tokens;
// remember the frist and last batch index for each sequence
// remember the first and last batch index for each sequence
std::fill(i_batch_beg.begin(), i_batch_beg.end(), -1);
std::fill(i_batch_end.begin(), i_batch_end.end(), -1);
@@ -1621,7 +1621,7 @@ 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.n_past, { seq_id }, true);
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);
i_last[seq_id] = batch.n_tokens - 1;
@@ -1635,16 +1635,16 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
while (n_drafting > 0) {
// each step decodes under a different head, i.e. a different decoder layer, and
// KV is per layer. process() filled this layer's KV only for positions < n_past
// KV is per layer. process() filled this layer's KV only for positions < pos0
// (prompt + accepted prefix) — nothing in the draft region yet. so reset the
// draft region (the seq_rm lower bound is n_past, leaving the prompt KV intact)
// draft region (the seq_rm lower bound is pos0, leaving the prompt KV intact)
// and select head i so it rebuilds its own layer's KV there; decoding just the
// latest token would leave its attention reading cells only another head wrote.
if (chain_heads) {
auto * mem_dft = llama_get_memory(ctx_dft);
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (drafting[seq_id]) {
llama_memory_seq_rm(mem_dft, seq_id, dparams[seq_id].n_past, -1);
llama_memory_seq_rm(mem_dft, seq_id, dparams[seq_id].pos0, -1);
}
}
llama_set_nextn_layer_offset(ctx_dft, i);
@@ -1710,17 +1710,17 @@ 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.n_past + t, { seq_id }, t == n_rows - 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);
}
} 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.n_past, { seq_id }, true);
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);
} else {
common_batch_add(batch, id, dp.n_past + i + 1, { seq_id }, true);
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);
}
+1 -1
View File
@@ -61,7 +61,7 @@ struct common_speculative_draft_params {
// can be used to constraint the max draft based on the remaining context size
int32_t n_max = -1;
llama_pos n_past;
llama_pos pos0;
llama_token id_last;
// TODO: remove in the future by keeping track of the prompt from the _begin() call and the consecutive accept calls
+119 -22
View File
@@ -188,7 +188,7 @@ llama_memory_breakdown_print: | - Host | 439 =
Op test for MUL_MAT:
```
~/src/llama.cpp$ ./scripts/snapdragon/run.py --target adb --hex-hostbuf 0 --devices HTP0:0 -- test-backend-ops -b HTP0:0 -o MUL_MAT
~/src/llama.cpp$ ./scripts/snapdragon/run.py --target adb --devices HTP0:0 -- test-backend-ops -b HTP0:0 -o MUL_MAT
...
Backend 2/3: HTP0:0
Device description: Hexagon
@@ -213,14 +213,109 @@ ggml-hex: new session: HTP0 : session-id 0 domain-id 3 uri file:///libggml-htp-v
| llama 1B Q4_0 | 729.75 MiB | 1.24 B | HTP | 99 | 4 | 128 | 0 | tg64 | 51.54 ± 1.13 |
```
## Multi-Device Execution Modes
The Hexagon backend supports multiple execution and partitioning modes to accommodate different model sizes, memory
constraints, and single- or multi-NPU hardware topologies:
### 1. Single-Device Mode with Dynamic Buffer Mapping
Runs the model on a single NPU session (e.g. `HTP0` or `HTP0:0`).
A single NPU session provides ~3.5GB of available virtual address space. For models larger than 3.5GB, the backend
automatically maps and unmaps weight buffers during graph execution. This allows large models to run on a single NPU
without manual configuration:
```bash
./scripts/snapdragon/run.py --target adb --devices HTP0:0 -- \
llama-cli -m models/Llama-3.2-3B-Instruct-Q4_0.gguf -ngl 99 -p "Hello"
```
### 2. Layer-Split Mode across Virtual Sessions (`HTP0,HTP1,...` or `HTP0:0,HTP0:1,...`)
Partitions model layers at load time across multiple virtual sessions hosted on a single physical NPU.
Each virtual session acts as an independent backend device from llama.cpp's perspective (similar to multiple GPUs).
Because layers are permanently distributed across sessions, each session's allocated weights remain within its private 3.5GB
address space window, eliminating runtime buffer re-mapping overhead.
Here is an example of running the GPT-OSS-20B model on a Snapdragon device using 4 virtual sessions on a single NPU:
```bash
./scripts/snapdragon/run.py --target adb \
--devices HTP0:0,HTP0:1,HTP0:2,HTP0:3 -- \
llama-cli --load-mode none -m /data/local/tmp/gguf/gpt-oss-20b-Q4_0.gguf -t 4 \
--ctx-size 8192 --batch-size 128 -ctk q8_0 -ctv q8_0 -fa on -ngl 99 -no-cnv -f surfing.txt
```
Log output snippet:
```
...
llama_model_loader: - type f32: 289 tensors
llama_model_loader: - type q4_0: 96 tensors
llama_model_loader: - type q8_0: 2 tensors
llama_model_loader: - type mxfp4: 72 tensors
...
load_tensors: offloaded 25/25 layers to GPU
load_tensors: CPU model buffer size = 1182.09 MiB
load_tensors: HTP0:1 model buffer size = 2512.58 MiB
load_tensors: HTP0:3 model buffer size = 2093.83 MiB
load_tensors: HTP0:0 model buffer size = 2931.34 MiB
load_tensors: HTP0:2 model buffer size = 2512.58 MiB
...
llama_perf_context_print: prompt eval time = 3843.67 ms / 197 tokens ( 19.51 ms per token, 51.25 tokens per second)
llama_perf_context_print: eval time = 1686.13 ms / 31 runs ( 54.39 ms per token, 18.39 tokens per second)
llama_perf_context_print: total time = 6266.30 ms / 228 tokens
llama_memory_breakdown_print: | memory breakdown [MiB] | total free self model context compute unaccounted |
llama_memory_breakdown_print: | - HTP0:0 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - HTP0:1 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - HTP0:2 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - HTP0:3 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - Host | 1476 = 1208 + 105 + 162 |
```
### 3. Tensor-Split Mode across Physical Devices (`HTP0:0,HTP1:0,...`)
Distributes model tensors across distinct physical NPU hardware cores using llama.cpp's tensor parallelism
(`--split-mode tensor`).
Tensors are partitioned across physical NPUs for parallel execution (proportions are distributed equally by default without
needing an explicit `--tensor-split` option):
```bash
./scripts/snapdragon/run.py --target adb \
--devices HTP0:0,HTP1:0 -- \
llama-cli -m models/Llama-3.2-3B-Instruct-Q4_0.gguf --split-mode tensor -ngl 99 -p "Hello"
```
### 4. Row-Split Multi-Device Mode via Device Grouping (`HTP0[0-1]`)
Groups multiple physical NPU cores into a single logical device using bracket notation (`HTP0[0-1]` or `HTP0[0,1]`).
Unlike host-level tensor-splitting, row-splitting is executed entirely inside the Hexagon backend:
```bash
./scripts/snapdragon/run.py --target adb \
--devices 'HTP0[0-1]' -- \
llama-cli -m models/Llama-3.2-3B-Instruct-Q4_0.gguf -ngl 99 -p "Hello"
```
You can also combine row-splitting with layer-splitting across multiple grouped devices (e.g. `--devices 'HTP0[0-1],HTP1[2-3]'`
on 4 physical NPUs, or `--devices 'HTP0[0-1:0],HTP1[0-1:1]'` on 2 physical NPUs using virtual sessions 0 and 1).
## Environment variables
- `GGML_HEXAGON_DEVICES` (default: not set, defaults to HTP0 session)
Controls which NPU devices and sessions to allocate. Can be configured as:
- A single integer `N`: Allocates `N` sessions named `HTP0`, `HTP1`, ..., `HTP<N-1>` (behaves identically to `GGML_HEXAGON_NDEV=N`).
- A comma-separated list of device names in `HTP<physical_idx>:<virtual_idx>` format (or legacy `HTP<idx>` format). For example, `HTP0:0,HTP0:1` creates two virtual
sessions on the first physical NPU (useful for memory limits). `HTP0:0,HTP1:0` allocates one session on each of the two physical NPUs
on a dual-NPU device.
Controls which NPU devices and sessions to allocate. Configurable via `--devices` in `run.py`:
- `N` (single integer): Allocates `N` virtual sessions named `HTP0`, `HTP1`, ..., `HTP<N-1>` on physical NPU 0.
- `HTP<phys>:<virt>,...`: Comma-separated list of individual devices specifying physical and virtual index:
- `HTP0:0,HTP0:1`: Two virtual sessions on physical NPU 0 (layer-split on single NPU).
- `HTP0:0,HTP1:0`: One session on physical NPU 0 and one on physical NPU 1 (tensor-split across physical cores).
- `HTP<name>[<phys_spec>]`: Device grouping syntax for row-split multi-device execution:
- `HTP0[0-1]`: A single logical device `HTP0` that groups physical cores 0 and 1.
- `HTP0[0-1],HTP1[2-3]`: Two layer-split devices across 4 physical NPUs (cores 0-1 and 2-3).
- `HTP0[0-1:0],HTP1[0-1:1]`: Two layer-split devices across 2 physical NPUs using virtual sessions 0 and 1.
- `GGML_HEXAGON_NDEV` (deprecated)
Replaced by `GGML_HEXAGON_DEVICES`. Controls the number of virtual sessions to allocate on physical NPU `0`.
@@ -229,9 +324,8 @@ ggml-hex: new session: HTP0 : session-id 0 domain-id 3 uri file:///libggml-htp-v
- `GGML_HEXAGON_NHVX=0`
Controls the number of HVX hardware threads to use. The default is all (actual number varies depending on the hardware version).
- `GGML_HEXAGON_HOSTBUF=1`
Controls whether the Hexagon backend allocates host buffers. By default, all buffers except for REPACK are host buffers.
This option is required for testing Ops that require REPACK buffers (MUL_MAT and MUL_MAT_ID).
- `GGML_HEXAGON_HOSTBUF=1` (default: 0, disabled)
Enables allocating host buffers for debugging. By default, host buffers are disabled.
- `GGML_HEXAGON_VERBOSE=1`
Enables verbose logging of Ops from the backend. Example output:
@@ -246,23 +340,26 @@ ggml-hex: new session: HTP0 : session-id 0 domain-id 3 uri file:///libggml-htp-v
```
- `GGML_HEXAGON_PROFILE=1`
Enables Op profiling:
Enables Op profiling (configurable via `--hex-profile` in `run.py`):
- `1` Basic profile with per-op `usecs` and `cycles` counters
- `2` Extended profile with per-op `usecs`, `cycles` and default PMU counter data
- `0x1,...,0x8` Extended profile with per-op `usecs`, `cycles` and custom PMU counter data
- `1`: Basic profile with per-op `usecs` and `cycles` counters
- `2`: Extended profile with per-op `usecs`, `cycles` and default PMU counter data
- `0x1,...,0x8`: Extended profile with per-op `usecs`, `cycles` and custom PMU counter data
The logging output can be either saved into a file for post-processing or it can be piped directly into the post-processing tool
to generate the report.
Examples:
The logging output can be saved to a file or piped directly into the post-processing script:
`GGML_HEXAGON_PROFILE=1 ./scripts/snapdragon/run.py --target adb -- llama-cli ... |& ./scripts/snapdragon/ggml-hexagon-profile.py -`
```bash
./scripts/snapdragon/run.py --target adb --hex-profile 1 -- llama-cli ... |& \
./scripts/snapdragon/ggml-hexagon-profile.py -
```
- `GGML_HEXAGON_OPFILTER=regex`
Allows filtering (disabling) Ops that match the regex pattern:
Filters (disables) Ops matching the regex pattern (configurable via `--hex-opfilter` in `run.py`):
Examples:
`GGML_HEXAGON_OPFILTER="FLASH_ATTN_EXT" ./scripts/snapdragon/run.py --target adb -- llama-cli ...` - Disable Flash Attention on Hexagon (falls back to CPU or GPU)
`GGML_HEXAGON_OPFILTER="ADD\|SUB" ./scripts/snapdragon/run.py --target adb -- llama-cli ...` - Disable ADD and SUB on Hexagon (fall back to CPU or GPU)
```bash
# Disable Flash Attention on Hexagon (falls back to CPU or GPU)
./scripts/snapdragon/run.py --target adb --hex-opfilter "FLASH_ATTN_EXT" -- llama-cli ...
# Disable ADD and SUB on Hexagon (fall back to CPU or GPU)
./scripts/snapdragon/run.py --target adb --hex-opfilter "ADD|SUB" -- llama-cli ...
```
+296 -64
View File
@@ -2,16 +2,16 @@
## Backend libraries
The Hexagon backend consist of two parts:
The Hexagon backend consists of two parts:
- `libggml-hexagon`
This is the regular CPU-side GGML backend library, either shared or statically linked
This is the regular CPU-side GGML backend library, either shared or statically linked.
- `libggml-htp-vNN`
This is the NPU-side (HTP stands for Hexagon Tensor Processor) shared library that contains the Op dispatcher and kernels.
The correct library is selected automatically at runtime based on the HW version.
Here is an example of the build artifacts
Here is an example of the build artifacts:
```
~/src/llama.cpp$ ls -l pkg-adb/llama.cpp/lib/libggml*
@@ -26,75 +26,307 @@ pkg-adb/llama.cpp/lib/libggml-htp-v81.so
## Memory buffers
Hexagon NPU backend takes advantage of the Snapdragon's unified memory model where all buffers are fully accessible by the CPU and GPU.
The NPU does have a dedicated tightly-coupled memory called VTCM but that memory is used only for intermediate data (e.g. dynamically
quantized tensors) or temporary data (chunks of the weight tensors fetched via DMA).
Please note that currently the Hexagon backend does not implement SET/GET_ROWS Ops because there is no advantage in offloading those
to the NPU at this point.
The backend does allocates non-host buffers for the tensors with datatypes that require repacking: Q4_0, Q8_0, MXFP4.
From the MMU perspective these buffers are still regular buffers (normal access by the CPU) they are marked as non-host simply to force
the repacking.
The Hexagon NPU backend takes advantage of Snapdragon unified memory where all DDR buffers are accessible by CPU, GPU, and NPU.
The NPU has dedicated tightly-coupled memory called VTCM (Vector Tightly-Coupled Memory). VTCM is used for intermediate data (such as
dynamically quantized activations) and streaming buffers (chunks of weight and activation tensors fetched via DMA).
## Large model handling
Hexagon NPU sessions (aka Process Domains (PD) in the Hexagon SDK) are limited to a maximum memory mapping window of around 3.5GB.
Hexagon NPU sessions have a 32-bit virtual address space window of around 3.5GB.
In llama.cpp/GGML, each Hexagon session is mapped to a single GGML backend device (e.g., `HTP0:0`, `HTP0:1`, etc. when using
`GGML_HEXAGON_DEVICES`, or `HTP0`, `HTP1` in legacy mode).
To support running models larger than 3.5GB on a single device, the Hexagon backend dynamically maps and unmaps execution buffers
during the graph execution cycle to stay within the Process Domain window. This enables large models to run successfully on a single
NPU device.
To support running models larger than 3.5GB on a single device, the Hexagon backend dynamically maps and unmaps buffers:
- Buffers are allocated in shared DDR (RPCMEM) via file descriptors (`fastrpc_mmap` using `FASTRPC_MAP_FD_DELAYED`).
- Pinned buffers (such as KV cache and active compute buffers) remain mapped throughout execution.
- Inactive weight buffers are dynamically mapped into the NPU session via `HAP_mmap()` during batch buffer preparation
(`prep_op_bufs()` in `htp/main.c`) and unmapped via `htp_iface_munmap()` when no longer needed by the active batch.
- This dynamic sliding window allows a single NPU session to execute models that exceed the 3.5GB window.
Alternatively, users can choose to use standard llama.cpp/GGML layer-splitting mode to partition and split the model across
multiple Hexagon devices or virtual sessions (which behave like multiple GPUs from the offload and splitting perspective).
Alternatively, users can partition and split the model across multiple virtual sessions or physical NPUs using layer-splitting,
tensor-splitting, or row-splitting modes. For user-facing execution modes and examples, see the
[Snapdragon user guide](README.md#multi-device-execution-modes).
Here is an example of running GPT-OSS-20B model on a Snapdragon device using 4 virtual sessions on a single NPU (physical index 0).
## Op and Kernel Development Guidelines
Writing high-performance operators for Hexagon requires following specific guidelines.
### DDR -> DMA -> VTCM Execution Pipeline
- Strongly prefer the `DDR -> DMA -> VTCM -> compute (HVX/HMX) -> VTCM -> DMA -> DDR` data flow.
- Direct HVX reads/writes from/to DDR are less efficient and should only be used as a fallback.
- The DMA queue is a strict FIFO where operations must be pushed and popped in strict order.
- Follow the pipelined multi-buffering sequence properly (typically 2x to 16x buffering) so every push has a corresponding pop:
1. In the prologue, push initial DDR -> VTCM transfers to prime the pipeline.
2. In the loop body, wait for buffer N via DMA pop, launch HVX/HMX compute on buffer N, push VTCM -> DDR writeback of result N,
and push DDR -> VTCM prefetch of buffer N+2.
3. In the epilogue, pop all remaining in-flight transfers to drain the pipeline.
- Because every push must be matched by a pop, `dma_queue_flush()` is not required when the pipeline sequence is followed
properly. Flushing is only used in rare exceptions where a batch of operations is pushed without individual pops.
- Use the DMA queue interface from [`dma-queue.h`](../../../ggml/src/ggml-hexagon/htp/dma-queue.h)
(`dma_queue_push_ddr_to_vtcm`, `dma_queue_pop`, `dma_queue_push_vtcm_to_ddr`).
See [`cumsum-ops.c`](../../../ggml/src/ggml-hexagon/htp/cumsum-ops.c) and
[`act-ops.c`](../../../ggml/src/ggml-hexagon/htp/act-ops.c) for reference implementations.
### Avoid Scalar Reads and Writes to VTCM
- Access VTCM data using DMA transfers or HVX/HMX vector instructions rather than scalar reads and writes.
### Avoid Scalar Division in Inner Loops
- Hexagon cores do not have hardware division instructions.
- For recurring divisions across iterations or threads, use `fastdiv` from
[`hex-fastdiv.h`](../../../ggml/src/ggml-hexagon/htp/hex-fastdiv.h) with precomputed divisors (such as
`octx->ctx->mdev.count_div` or `octx->n_threads_div`).
- Do not call `init_fastdiv_values()` for single-use divisions; use standard compiler division (`/`) instead.
### Host-Side Precomputation via `kernel_params`
- Precompute tensor shapes, strides, scale conversions, tiling layouts, and validation checks on the host CPU during graph
preparation in [`ggml-hexagon.cpp`](../../../ggml/src/ggml-hexagon/ggml-hexagon.cpp).
- Pack precomputed parameters into the operator's fixed `kernel_params` structure in `htp_op_node` (such as
`htp_mm_kernel_params`, `htp_unary_kernel_params`, `htp_fa_kernel_params`, `htp_get_rows_kernel_params`).
- The NPU executes directly using `octx->kernel_params` without redundant runtime metadata extraction or validation.
- **Strict Host-Kernel Alignment**:
- Verify that parameters calculated by the host CPU are strictly honored by the NPU kernel.
- Ensure the kernel does not ignore host-computed fields (for example, falling back to `octx->n_threads` instead of
using `kparams->n_threads`, or ignoring precomputed `tasks_per_thread` and chunk counts).
- Both human developers and coding agents must audit both sides of the interface: ensure fields populated in `kernel_params`
in [`ggml-hexagon.cpp`](../../../ggml/src/ggml-hexagon/ggml-hexagon.cpp) are actively and consistently utilized by the
corresponding operator entry point and worker threads in `htp/*-ops.c`.
### Tracing Instrumentation
- All kernels must include trace events for performance profiling and timeline visualization in Perfetto
([`hex-profile.h`](../../../ggml/src/ggml-hexagon/htp/hex-profile.h)).
- Surround compute sections with `htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) info)` and
`htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) info)`.
- Use specific event types for major phases:
- `HTP_TRACE_EVT_HVX_COMP`: Vector compute execution.
- `HTP_TRACE_EVT_DMA`: DMA transfer wait or poll cycles.
- `HTP_TRACE_EVT_FENCE`: Multi-device fence barrier synchronization.
- `HTP_TRACE_EVT_L2FLUSH`: L2 cache cleaning operations.
- Pass meaningful progress metrics (such as row index, chunk index, or token index) in the 16-bit `info` parameter.
### Work Queue and Threading
- Distribute parallel work across NPU worker threads using the thread pool work queue:
```c
work_queue_run(ctx->work_queue, worker_func, &op_ctx, n_threads);
```
- Keep worker functions independent and re-entrant. Worker threads should only operate on their designated chunk of rows or elements.
### Avoid Redundant Defensive NULL Checks
- Do not add defensive NULL checks or assertions for internal framework pointers or required graph operands and outputs.
Internal pointers include `ctx`, `octx`, local context structs like `*ctx`, `kparams`, and worker callback `data`.
- These pointers are architectural invariants during kernel execution and host-side graph preparation.
Graph compute receives allocated nodes with valid required `node->src[N]` and `node->data` pointers.
- Do not turn an invariant violation into an unsupported operation or missed fusion.
Checks such as `if (!octx || !octx->ctx)` clutter the code, obscure intent, and hide upstream errors.
- **Distinction**: `octx->src[N]` pointers *can* be NULL by design and must be checked when optional.
Examples include attention masks, optional bias or weights in fused kernels, and frequency factors.
### Multiline Macro Formatting
- Keep trailing backslashes in multiline `#define` macros cleanly aligned to a consistent column.
- Avoid trailing whitespace after macro backslashes.
- Use [`scripts/snapdragon/ggml-hexagon-align-macros.py`](../../../scripts/snapdragon/ggml-hexagon-align-macros.py) to inspect, diff,
or automatically align macro definitions across Hexagon kernel sources:
```bash
# Check for misaligned macros
python3 scripts/snapdragon/ggml-hexagon-align-macros.py ggml/src/ggml-hexagon/htp/
# Fix misaligned macros in-place
python3 scripts/snapdragon/ggml-hexagon-align-macros.py --fix ggml/src/ggml-hexagon/htp/
```
## Multi-Device Partitioning (mdev)
Multi-device (mdev) mode enables row-level tensor parallel execution across multiple physical NPU cores or virtual NPU
sessions.
### 128-Byte Cache Line Alignment
- Shared tensor buffers reside in DDR (RPCMEM) with a 128-byte cache line granularity
(`HEX_L2_LINE_SIZE` = 128 bytes, `HTP_TENSOR_MDEV_LINE_SIZE`).
- **Rule**: Multi-device work partitions must align destination write regions to 128-byte cache line boundaries so distinct
devices never share or overwrite the same cache line.
### Partitioning Helpers in `htp-tensor.h`
Common partitioning logic is factored into reusable inline helpers in
[`htp-tensor.h`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h):
1. [`htp_tensor_mdev_rows_per_chunk`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L67):
Determines the minimum number of rows per chunk so that the chunk byte size is a multiple of 128 bytes:
```
rows_per_chunk = 128 / hex_gcd_u32(row_size, 128)
```
If row stride `nb[1]` is already a multiple of 128 bytes, `rows_per_chunk = 1`.
Returns `false` if the tensor cannot be safely row-partitioned (such as unaligned base pointer, permuted layout,
or non-128-byte aligned outer strides).
2. [`htp_tensor_mdev_partition`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L94):
Calculates the per-device work range `struct htp_tensor_mdev_range { uint32_t start; uint32_t count; }` given
`total_units`, `units_per_chunk`, `mdev_idx`, `mdev_count`, and the precomputed `mdev_count_div`.
Handles chunk distribution across devices, assigns remainder units to the last device, and automatically triggers
single-device fallback when partitioning is unsafe.
### Row-Partitioned Operators
For row-wise operators
(such as activations in [`act-ops.c`](../../../ggml/src/ggml-hexagon/htp/act-ops.c),
binary ops in [`binary-ops.c`](../../../ggml/src/ggml-hexagon/htp/binary-ops.c),
unary ops in [`unary-ops.c`](../../../ggml/src/ggml-hexagon/htp/unary-ops.c), and
sameshape copies in [`cpy-ops.c`](../../../ggml/src/ggml-hexagon/htp/cpy-ops.c)):
```c
const uint32_t total_rows = ne01 * ne02 * ne03;
const size_t dst_row_size = dst->ne[0] * elem_size;
uint32_t row_start = 0;
uint32_t nrows = total_rows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
total_rows, rows_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;
}
```
### Element-Partitioned Operators
For flat element-wise operations (such as reshape copies in
[`cpy-ops.c`](../../../ggml/src/ggml-hexagon/htp/cpy-ops.c)):
- Partition total linear elements N = ne0 * ne1 * ne2 * ne3 in 128-byte cache line chunks (`elems_per_line = (elem_size == 4) ? 32 : 64`).
- Requires strict 1D contiguity:
[`htp_tensor_is_contiguous(dst, elem_size)`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L28)
and 128-byte aligned destination pointer
[`htp_tensor_mdev_data_aligned(dst)`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L47).
- If contiguous and aligned, pass `elems_per_line` to
[`htp_tensor_mdev_partition`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h#L94);
otherwise pass 0 to trigger Device 0 fallback.
### Single-Device Fallback (Device 0)
- Fallback to Device 0 (`mdev.idx == 0`) when partitioning would cause cache line tearing or when work cannot be evenly distributed.
- Triggers:
1. Destination tensor cannot be safely partitioned (`rows_per_chunk == 0` or non-contiguous/unaligned buffer).
2. Total aligned chunks < `mdev_count`.
- Device 0 processes the entire tensor `[0, total_units)`.
- Devices 1 ... N-1 receive `count = 0` and return `HTP_STATUS_OK` immediately.
### Flatten Outer Dimensions Globally
- **Never partition solely on `ne01` (dimension 1).**
- Partitioning only on `ne01` repeats the device boundary across every 2D slice (`ne02`, `ne03`). If each 2D slice is small,
false sharing occurs repeatedly throughout the tensor.
- Always flatten outer dimensions globally: `total_rows = ne01 * ne02 * ne03` and partition once across the combined row space.
### Stateless Starting Coordinates
- Do not use incremental state variables across slices that assume the thread or device starts at index 0.
- Precompute starting multidimensional coordinates at `r = row_start` (or `e = elem_start`) once using `fastdiv`.
- In inner loops, step base pointers directly (`ptr += stride`) or reset/wrap coordinates explicitly (`if (++i01 == ne01) { ... }`).
### Clean Range Encapsulation
- Initialize single-device default ranges at declaration:
```c
uint32_t row_start = 0;
uint32_t nrows = total_rows;
```
- Encapsulate all multi-device logic inside `if (octx->ctx->mdev.count > 1)`. If the block is omitted or compiled out,
the operator runs standard single-device execution untouched.
- Do not propagate `mdev_` prefixes to worker functions or context structs. Worker threads are device-agnostic and
should only receive standard range parameters (`ctx.row_start`, `ctx.nrows`).
- In worker threads, calculate row intervals using standard arithmetic:
```c
const uint32_t ir0 = ctx->row_start + dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, ctx->row_start + ctx->nrows);
```
In single-device mode (`row_start == 0`), this naturally simplifies to `dr * ith` and `MIN(ir0 + dr, ctx->nrows)` with zero overhead.
## Multi-Device Synchronization
Multi-device execution synchronizes worker sessions across devices using explicit barriers and tensor cache flushing.
### Synchronization Fence Protocol
Multi-device execution synchronizes worker sessions through atomic fence slots and barriers defined in
[`htp-fence.h`](../../../ggml/src/ggml-hexagon/htp/htp-fence.h):
```
~/src/llama.cpp$ ./scripts/snapdragon/run.py --target adb --devices HTP0:0,HTP0:1,HTP0:2,HTP0:3 -- llama-cli --load-mode none -m /data/local/tmp/gguf/gpt-oss-20b-Q4_0.gguf -t 4 --ctx-size 8192 --batch-size 128 -ctk q8_0 -ctv q8_0 -fa on -ngl 99 -no-cnv -f surfing.txt
...
llama_model_loader: - type f32: 289 tensors
llama_model_loader: - type q4_0: 96 tensors
llama_model_loader: - type q8_0: 2 tensors
llama_model_loader: - type mxfp4: 72 tensors
...
load_tensors: offloaded 25/25 layers to GPU
load_tensors: CPU model buffer size = 1182.09 MiB
load_tensors: HTP0:1 model buffer size = 2512.58 MiB
load_tensors: HTP0:3 model buffer size = 2093.83 MiB
load_tensors: HTP0:0 model buffer size = 2931.34 MiB
load_tensors: HTP0:2 model buffer size = 2512.58 MiB
...
llama_context: n_ctx_per_seq (8192) < n_ctx_train (131072) -- the full capacity of the model will not be utilized
llama_context: CPU output buffer size = 0.77 MiB
llama_kv_cache_iswa: creating non-SWA KV cache, size = 8192 cells
llama_kv_cache: HTP0:1 KV buffer size = 25.50 MiB
llama_kv_cache: HTP0:3 KV buffer size = 25.50 MiB
llama_kv_cache: HTP0:0 KV buffer size = 25.50 MiB
llama_kv_cache: HTP0:2 KV buffer size = 25.50 MiB
llama_kv_cache: size = 102.00 MiB ( 8192 cells, 12 layers, 1/1 seqs), K (q8_0): 51.00 MiB, V (q8_0): 51.00 MiB
llama_kv_cache_iswa: creating SWA KV cache, size = 256 cells
llama_kv_cache: HTP0:1 KV buffer size = 0.80 MiB
llama_kv_cache: HTP0:3 KV buffer size = 0.53 MiB
llama_kv_cache: HTP0:0 KV buffer size = 1.06 MiB
llama_kv_cache: HTP0:2 KV buffer size = 0.80 MiB
llama_kv_cache: size = 3.19 MiB ( 256 cells, 12 layers, 1/1 seqs), K (q8_0): 1.59 MiB, V (q8_0): 1.59 MiB
llama_context: HTP0:0 compute buffer size = 16.06 MiB
llama_context: HTP0:1 compute buffer size = 16.06 MiB
llama_context: HTP0:2 compute buffer size = 16.06 MiB
llama_context: HTP0:3 compute buffer size = 16.06 MiB
llama_context: CPU compute buffer size = 98.19 MiB
...
llama_perf_context_print: prompt eval time = 3843.67 ms / 197 tokens ( 19.51 ms per token, 51.25 tokens per second)
llama_perf_context_print: eval time = 1686.13 ms / 31 runs ( 54.39 ms per token, 18.39 tokens per second)
llama_perf_context_print: total time = 6266.30 ms / 228 tokens
llama_perf_context_print: graphs reused = 30
llama_memory_breakdown_print: | memory breakdown [MiB] | total free self model context compute unaccounted |
llama_memory_breakdown_print: | - HTP0:0 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - HTP0:1 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - HTP0:2 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - HTP0:3 (Hexagon) | 2048 = 2048 + ( 0 = 0 + 0 + 0) + 0 |
llama_memory_breakdown_print: | - Host | 1476 = 1208 + 105 + 162 |
[NPU Session 0] [NPU Session 1]
| |
(Input Prep) (Input Prep)
| |
Pre-Op Barrier ----------------------------- Pre-Op Barrier
(mdev_sync_fence) (mdev_sync_fence)
| |
Kernel Execution Kernel Execution
(Output Slice 0) (Output Slice 1)
| |
Tensor Cache Flush Tensor Cache Flush
(htp_tensor_flush_all) (htp_tensor_flush_all)
| |
Post-Op/Batch Barrier ---------------------- Post-Op/Batch Barrier
(htp_mdev_group_barrier) (htp_mdev_group_barrier)
| |
Return Response to Host Return Response to Host
```
### Atomic Fence Slots and Cache Invalidation
- Fence synchronization operates on dedicated RPCMEM shared memory mapped across all participating sessions (`ctx->mdev.fence_base`).
- Each device owns a dedicated 128-byte cache-line aligned fence slot:
```c
atomic_uint * my_fence = htp_mdev_fence_slot(fence_base, mdev_idx);
```
- **Writing to fence ([`htp_fence_write`](../../../ggml/src/ggml-hexagon/htp/htp-fence.h#L18))**:
Stores `seq` and `status`, issues a `syncht` thread synchronization barrier, and flushes/invalidates the line
using `Q6_dccleaninva_A(fence)`.
- **Reading from peer fence ([`htp_fence_read`](../../../ggml/src/ggml-hexagon/htp/htp-fence.h#L26))**:
Executes `Q6_dccleaninva_A(fence)` and `syncht` before reading atomic values to ensure fresh data from DDR.
### Deterministic Monotonic Sequence Numbers
- Barrier fences use monotonically increasing sequence numbers:
```c
const uint32_t seq = ++ctx->mdev.fence_seq;
```
- Comparing sequence numbers with signed arithmetic `(int32_t)(peer_seq - seq) >= 0` prevents race conditions or
misaligned barrier arrivals across iterations.
- If any peer reports an error status (`peer_status > HTP_STATUS_OK`), the barrier propagates the error and unblocks immediately.
### Tensor Cache Flush and Pipeline Completion
- In the kernel, ensure all pushed DMA operations have been popped in strict FIFO order to drain the queue.
- Use [`htp_tensor_flush_all()`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h) to flush specific dirty tensors back to DDR:
- [`htp_tensor_flush_all()`](../../../ggml/src/ggml-hexagon/htp/htp-tensor.h) flushes only modified tensor address ranges,
ensuring peer devices and the host CPU observe consistent data in DDR.
- Never signal completion before all DMA transfers are drained and dirty tensor flushes have completed.
+122
View File
@@ -0,0 +1,122 @@
## Build profiling
This page is a working document for analyzing the current build and try to
identify ways to improve the build time.
### Requirements
The profiling script requires clang to be used as the compiler tool chain and
also requires that ClangBuildAnalyzer is installed.
Mac:
```console
brew install clang-build-analyzer
```
Linux:
```console
git clone https://github.com/aras-p/ClangBuildAnalyzer.git
cd ClangBuildAnalyzer
cmake -B build -DCMAKE_BUILD_TYPE=Release
cmake --build build -j$(nproc)
sudo cp build/ClangBuildAnalyzer /usr/local/bin/
```
Windows: install LLVM/clang and Ninja (e.g. via the
[LLVM releases page](https://github.com/llvm/llvm-project/releases) and
`winget install Ninja-build.Ninja`), then build ClangBuildAnalyzer the same
way as on Linux:
```console
git clone https://github.com/aras-p/ClangBuildAnalyzer.git
cd ClangBuildAnalyzer
cmake -B build -G Ninja -DCMAKE_C_COMPILER=clang -DCMAKE_CXX_COMPILER=clang++ -DCMAKE_BUILD_TYPE=Release
cmake --build build --config Release
```
Then add `ClangBuildAnalyzer\build` to `PATH`.
### Usage
Mac/Linux:
```console
$ ./scripts/build-profile.sh
```
Windows:
```console
> .\scripts\build-profile.ps1
```
Both accept `--full`/`-Full` (include Server, Tools, and Tests) and a jobs
override (`-jN` / `-Jobs N`).
Note: on Windows, `cmake` defaults to the Visual Studio generator, which
ignores `CMAKE_C_COMPILER`/`CMAKE_CXX_COMPILER` and silently falls back to
MSVC. `build-profile.ps1` passes `-G Ninja` so clang is actually used, this
is required on ARM64.
### Linux (Ubuntu 24.04)
Environment:
- Clang: 18.1.3 (Ubuntu clang version 18.1.3 (1ubuntu1))
- libstdc++: GCC 13.3.0 (Ubuntu 13.3.0-6ubuntu2~24.04.1)
- Target: x86_64-pc-linux-gnu
```console
+------------------------+-----+------------+------------+------------+
| Build | TUs | Frontend | Backend | Total |
+------------------------+-----+------------+------------+------------+
| Minimal, master | 249 | 468.2 s | 270.3 s | 738.5 s |
| Minimal, with PCH | 253 | 177.1 s | 265.8 s | 442.9 s |
| Full, master | 396 | 811.0 s | 692.2 s | 1,503.2 s |
| Full, with PCH | 405 | 380.0 s | 664.7 s | 1,044.7 s |
| Full, with PCH + UB | 264 | 357.7 s | 635.7 s | 993.4 s |
+------------------------+-----+------------+------------+------------+
PCH = precompiled header.
Full = includes building Server, Tools, and Tests.
UB = unity build for models
```
Note that the number of translation units (TUs) increases when using precompiled
headers — each PCH target adds one extra TU for the precompilation step itself.
### Mac (Apple M3)
Environment:
- Clang: Apple clang version 17.0.0 (clang-1700.3.19.1)
- libc++: ships with Apple clang 17.0.0 (Xcode toolchain)
- Target: arm64-apple-macosx15.6
```console
+------------------------+-----+------------+------------+------------+
| Build | TUs | Frontend | Backend | Total |
+------------------------+-----+------------+------------+------------+
| Minimal, master | 256 | 154.5 s | 94.8 s | 249.3 s |
| Minimal, with PCH | 261 | 65.9 s | 90.0 s | 155.9 s |
| Full, master | 407 | 265.7 s | 209.7 s | 475.4 s |
| Full, with PCH | 414 | 154.6 s | 197.5 s | 352.1 s |
| Full, with PCH + UB | 274 | 143.0 s | 192.2 s | 335.2 s |
+------------------------+-----+------------+------------+------------+
PCH = precompiled header.
Full = includes building Server, Tools, and Tests.
UB = unity build for models
```
### Windows (ARM64)
Environment:
- Clang: clang version 22.1.8 (LLVM, `C:\Program Files\LLVM`)
- STL: MSVC STL (Visual Studio 2022 Build Tools 14.44.35207)
- Target: aarch64-pc-windows-msvc
```console
+------------------------+-----+------------+------------+------------+
| Build | TUs | Frontend | Backend | Total |
+------------------------+-----+------------+------------+------------+
| Minimal, master | 249 | 159.4 s | 82.2 s | 241.6 s |
| Full, master | 373 | 337.2 s | 167.4 s | 504.6 s |
| Minimal, with PCH + UB | 113 | 62.3 s | 82.4 s | 144.7 s |
| Full, with PCH + UB | 240 | 233.0 s | 185.1 s | 418.1 s |
+------------------------+-----+------------+------------+------------+
PCH = precompiled header.
Full = includes building Server, Tools, and Tests.
UB = unity build for models
```
+1 -1
View File
@@ -806,7 +806,7 @@ To read documentation for how to build on Android, [click here](./android.md)
## WebGPU
The WebGPU backend relies on [Dawn](https://dawn.googlesource.com/dawn). Follow the instructions [here](https://dawn.googlesource.com/dawn/+/refs/heads/main/docs/quickstart-cmake.md) to install Dawn locally so that llama.cpp can find it using CMake. The current implementation is up-to-date with Dawn commit `18eb229`.
The WebGPU backend relies on [Dawn](https://dawn.googlesource.com/dawn). Follow the instructions [here](https://dawn.googlesource.com/dawn/+/refs/heads/main/docs/quickstart-cmake.md) to install Dawn locally so that llama.cpp can find it using CMake. The current implementation is up-to-date with Dawn commit `94c3c9c`.
In the llama.cpp directory, build with CMake:
@@ -188,7 +188,7 @@ int main(int argc, char ** argv) {
common_speculative_get_draft_params(spec, seq_id) = {
/* .drafting = */ true,
/* .n_max = */ n_draft_max,
/* .n_past = */ n_past,
/* .pos0 = */ n_past,
/* .id_last = */ id_last,
/* .prompt = */ &prompt_tgt,
/* .result = */ &draft, // output
+4
View File
@@ -1705,6 +1705,10 @@ static enum ggml_status ggml_backend_sched_compute_splits(ggml_backend_sched_t s
ggml_tensor * ids_tensor = node->src[2];
ggml_backend_t ids_backend = split_backend;
if (ggml_nelements(ids_tensor) == 0) {
continue;
}
// if the ids tensor is also an input of the split, it may not have been copied yet to the split backend
// in that case, we use the original ids tensor
for (int i = input_id + 1; i < split->n_inputs; i++) {
+6
View File
@@ -675,6 +675,12 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
target_compile_options(${GGML_CPU_NAME} PRIVATE ${ARCH_FLAGS})
target_compile_definitions(${GGML_CPU_NAME} PRIVATE ${ARCH_DEFINITIONS})
if (CMAKE_C_COMPILER_ID STREQUAL "GNU" AND NOT GGML_SYSTEM_ARCH STREQUAL "x86")
message(STATUS "Skipping PCH for ${GGML_CPU_NAME}: GCC PCH is only enabled for x86 (arch: ${GGML_SYSTEM_ARCH})")
else()
target_precompile_headers(${GGML_CPU_NAME} PRIVATE ggml-impl.h)
endif()
if (EMSCRIPTEN)
set_target_properties(${GGML_CPU_NAME} PROPERTIES COMPILE_FLAGS "-msimd128")
endif()
+8
View File
@@ -18,7 +18,15 @@
#endif
#endif
// -Winterference-size was introduced in GCC 12
#if defined(__cplusplus) && defined(__GNUC__) && !defined(__clang__) && __GNUC__ >= 12
#pragma GCC diagnostic push
#pragma GCC diagnostic ignored "-Winterference-size"
#endif
static const size_t CACHE_LINE_SIZE_F32 = CACHE_LINE_SIZE/sizeof(float);
#if defined(__cplusplus) && defined(__GNUC__) && !defined(__clang__) && __GNUC__ >= 12
#pragma GCC diagnostic pop
#endif
// Work buffer size for im2col operations in CONV2D
#define GGML_IM2COL_WORK_SIZE (16 * 1024 * 1024)
+281
View File
@@ -0,0 +1,281 @@
static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_gcn(ggml_type type, int J, bool fallback) {
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q1_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q1_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q1_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_0, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_0, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_0, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_0, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_1, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_1, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_1, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_1, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_1, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_1, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_1, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_1, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_1, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 256, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 512, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q8_0, 512, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
// ---------------------------------------------------------------------------------------------
CASE(GGML_TYPE_Q2_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_K, 512, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_K, 512, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q2_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_K, 512, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_K, 512, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_K, 512, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_K, 512, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q2_K, 512, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q3_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q3_K, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q3_K, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q3_K, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q3_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q3_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q3_K, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q3_K, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q3_K, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q3_K, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q3_K, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_K, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q4_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_K, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_K, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_K, 256, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_K, 256, 2, 128, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q4_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_K, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q5_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_K, 256, 3, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_K, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_K, 256, 2, 128, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_K, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q5_K, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_Q6_K, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false);
// ---------------------------------------------------------------------------------------------
CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ1_S, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ1_S, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ1_S, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ1_S, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XXS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XXS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XXS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XXS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_XS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_S, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_S, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_S, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ2_S, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_XXS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_XXS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_XXS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_XXS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_S, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_S, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_S, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ3_S, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_XS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_XS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_XS, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_XS, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_NL, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_NL, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_NL, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_IQ4_NL, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false);
// ---------------------------------------------------------------------------------------------
CASE(GGML_TYPE_MXFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_MXFP4, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_MXFP4, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_MXFP4, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_MXFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_MXFP4, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_MXFP4, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_MXFP4, 512, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_MXFP4, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_MXFP4, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_MXFP4, 512, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_NVFP4, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_NVFP4, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_NVFP4, 256, 2, 128, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false);
CASE(GGML_TYPE_NVFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false);
return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true);
}
+7 -1
View File
@@ -218,6 +218,7 @@ struct ggml_cuda_mmq_config {
#include "mmq-config-ampere.cuh"
#include "mmq-config-blackwell.cuh"
#include "mmq-config-gcn.cuh"
#include "mmq-config-cdna.cuh"
#include "mmq-config-rdna2.cuh"
#include "mmq-config-rdna3.cuh"
@@ -228,6 +229,9 @@ struct ggml_cuda_mmq_config {
static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc) {
if (GGML_CUDA_CC_IS_AMD(cc)) {
if (GGML_CUDA_CC_IS_GCN(cc)) {
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
}
if (GGML_CUDA_CC_IS_CDNA(cc)) {
return ggml_cuda_mmq_get_config_cdna(type, J, fallback);
}
@@ -256,7 +260,9 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty
static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback) {
#ifdef GGML_USE_HIP
#ifdef CDNA
#ifdef GCN
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
#elif defined(CDNA)
return ggml_cuda_mmq_get_config_cdna(type, J, fallback);
#elif defined(RDNA4)
return ggml_cuda_mmq_get_config_rdna4(type, J, fallback);
File diff suppressed because it is too large Load Diff
+6
View File
@@ -344,6 +344,12 @@ struct htp_opformat {
} else if (htp_op_is_unary(node.opcode)) {
const auto * kparams = (const struct htp_unary_kernel_params *) node.kernel_params;
snprintf(str, max_size, "%s vtcm %d", kparams->col_tile ? "wide-row" : "row-block", (int) kparams->vtcm_size);
} else if (node.opcode == HTP_OP_MDEV_GROUP && node.node) {
snprintf(str, max_size, "idx %d count %d", (int) node.node->op_params[0], (int) node.dst()->ne[1]);
} else if ((node.opcode == HTP_OP_FENCE || node.opcode == HTP_OP_CPY_FENCE) && node.node) {
snprintf(str, max_size, "seq 0x%x", (uint32_t) node.node->op_params[0]);
} else if (node.opcode == HTP_OP_ALLREDUCE && node.node) {
snprintf(str, max_size, "seq 0x%x -> 0x%x", (uint32_t) node.node->op_params[0], (uint32_t) node.node->op_params[1]);
} else {
snprintf(str, max_size, "----");
}
+125 -108
View File
@@ -3,7 +3,6 @@
#pragma clang diagnostic ignored "-Wunused-but-set-variable"
#include <HAP_farf.h>
#include <HAP_perf.h>
#include <math.h>
#include <string.h>
@@ -15,7 +14,7 @@
#include "ggml-common.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-ops.h"
#include "hex-common.h"
#include "htp-tensor.h"
#include "htp-vtcm.h"
@@ -80,6 +79,7 @@ struct htp_act_context {
uint32_t block;
uint32_t src0_nrows;
uint32_t src0_nrows_per_thread;
uint32_t row_start;
int nc;
uint8_t * vtcm_src0;
@@ -329,104 +329,104 @@ static void geglu_f32(const float * restrict src0,
}
}
#define DEFINE_GLU_PER_THREAD(NAME, OP_STR, CORE_EXPR) \
static void glu_##NAME##_f32_per_thread(unsigned int nth, unsigned int ith, void * data) { \
struct htp_act_context * actx = (struct htp_act_context *) data; \
htp_act_preamble; \
\
struct htp_thread_trace * tr = actx->octx->ctx ? &actx->octx->ctx->trace[ith] : NULL; \
\
size_t src0_row_size = actx->src0_row_size; \
size_t src1_row_size = actx->src1_row_size; \
size_t dst_row_size = actx->dst_row_size; \
\
size_t src0_row_stride = actx->src0_row_stride; \
size_t src1_row_stride = actx->src1_row_stride; \
\
const uint32_t src0_nrows = actx->src0_nrows; \
const uint32_t src0_nrows_per_thread = actx->src0_nrows_per_thread; \
\
const uint32_t src0_start_row = src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); \
\
/* no work for this thread */ \
if (src0_start_row >= src0_end_row) { \
return; \
} \
\
const uint8_t * restrict data_src0 = actx->data_src0; \
const uint8_t * restrict data_src1 = actx->data_src1; \
uint8_t * restrict data_dst = actx->data_dst; \
\
const size_t src0_row_size_aligned = actx->src0_row_size_aligned; \
const size_t src1_row_size_aligned = actx->src1_row_size_aligned; \
const size_t dst_row_size_aligned = actx->dst_row_size_aligned; \
\
uint8_t * restrict src0_spad_data = actx->vtcm_src0 + (ith * actx->vtcm_src0_size_per_thread); \
uint8_t * restrict src1_spad_data = actx->vtcm_src1 + (ith * actx->vtcm_src1_size_per_thread); \
uint8_t * restrict dst_spad_data = actx->vtcm_dst + (ith * actx->vtcm_dst_size_per_thread); \
\
size_t src0_spad_half_size = actx->src0_spad_half_size; \
size_t src1_spad_half_size = actx->src1_spad_half_size; \
size_t dst_spad_half_size = actx->dst_spad_half_size; \
\
const int BLOCK = actx->block; \
if (BLOCK == 0) { \
FARF(ERROR, \
OP_STR \
" : current VTCM reservation %zu is too small for even 1 row per thread, needed at least %zu\n", \
actx->vtcm_src0_size_per_thread, src0_row_size_aligned); \
return; \
} \
\
dma_queue * dma_queue = actx->octx->ctx->dma[ith]; \
\
/* See discussion: https://github.com/ggml-org/llama.cpp/pull/18151#issuecomment-3678235379 */ \
for (uint32_t ir = src0_start_row, spad_idx = 0; ir < src0_end_row && spad_idx < 2; ir += BLOCK, spad_idx++) { \
const uint32_t block_size = MIN(BLOCK, src0_end_row - ir); \
\
/* Dummy DMA transation for sequencing (interleaving dst,src,dst,...) */ \
dma_queue_push_vtcm_to_ddr(dma_queue, \
dma_make_ptr(data_dst, dst_spad_data + (spad_idx * dst_spad_half_size)), \
dst_row_size, dst_row_size_aligned, 0); \
\
dma_queue_push( \
dma_queue, \
dma_make_ptr(src0_spad_data + (spad_idx * src0_spad_half_size), data_src0 + (ir * src0_row_stride)), \
src0_row_size_aligned, src0_row_stride, src0_row_size, block_size); \
dma_queue_push( \
dma_queue, \
dma_make_ptr(src1_spad_data + (spad_idx * src1_spad_half_size), data_src1 + (ir * src1_row_stride)), \
src1_row_size_aligned, src1_row_stride, src1_row_size, block_size); \
} \
\
for (uint32_t ir = src0_start_row; ir < src0_end_row; ir += BLOCK) { \
const uint32_t block_size = MIN(BLOCK, src0_end_row - ir); \
\
float * dst_spad = (float *) dma_queue_pop(dma_queue).src; \
float * src0_spad = (float *) dma_queue_pop(dma_queue).dst; \
float * src1_spad = (float *) dma_queue_pop(dma_queue).dst; \
\
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir); \
CORE_EXPR; \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir); \
\
dma_queue_push_vtcm_to_ddr(dma_queue, dma_make_ptr(data_dst + (ir * dst_row_size), dst_spad), \
dst_row_size, dst_row_size_aligned, block_size); \
\
/* prefetch N+2 loop iteration if any */ \
const uint32_t pref_block = (ir + BLOCK * 2); \
if (pref_block < src0_end_row) { \
const uint32_t pref_block_size = MIN(BLOCK, src0_end_row - pref_block); \
dma_queue_push(dma_queue, dma_make_ptr(src0_spad, data_src0 + (pref_block * src0_row_stride)), \
src0_row_size_aligned, src0_row_stride, src0_row_size, pref_block_size); \
dma_queue_push(dma_queue, dma_make_ptr(src1_spad, data_src1 + (pref_block * src1_row_stride)), \
src1_row_size_aligned, src1_row_stride, src1_row_size, pref_block_size); \
} \
} \
\
dma_queue_flush(dma_queue); \
\
#define DEFINE_GLU_PER_THREAD(NAME, OP_STR, CORE_EXPR) \
static void glu_##NAME##_f32_per_thread(unsigned int nth, unsigned int ith, void * data) { \
struct htp_act_context * actx = (struct htp_act_context *) data; \
htp_act_preamble; \
\
struct htp_thread_trace * tr = actx->octx->ctx ? &actx->octx->ctx->trace[ith] : NULL; \
\
size_t src0_row_size = actx->src0_row_size; \
size_t src1_row_size = actx->src1_row_size; \
size_t dst_row_size = actx->dst_row_size; \
\
size_t src0_row_stride = actx->src0_row_stride; \
size_t src1_row_stride = actx->src1_row_stride; \
\
const uint32_t src0_nrows = actx->src0_nrows; \
const uint32_t src0_nrows_per_thread = actx->src0_nrows_per_thread; \
\
const uint32_t src0_start_row = actx->row_start + src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, actx->row_start + src0_nrows); \
\
/* no work for this thread */ \
if (src0_start_row >= src0_end_row) { \
return; \
} \
\
const uint8_t * restrict data_src0 = actx->data_src0; \
const uint8_t * restrict data_src1 = actx->data_src1; \
uint8_t * restrict data_dst = actx->data_dst; \
\
const size_t src0_row_size_aligned = actx->src0_row_size_aligned; \
const size_t src1_row_size_aligned = actx->src1_row_size_aligned; \
const size_t dst_row_size_aligned = actx->dst_row_size_aligned; \
\
uint8_t * restrict src0_spad_data = actx->vtcm_src0 + (ith * actx->vtcm_src0_size_per_thread); \
uint8_t * restrict src1_spad_data = actx->vtcm_src1 + (ith * actx->vtcm_src1_size_per_thread); \
uint8_t * restrict dst_spad_data = actx->vtcm_dst + (ith * actx->vtcm_dst_size_per_thread); \
\
size_t src0_spad_half_size = actx->src0_spad_half_size; \
size_t src1_spad_half_size = actx->src1_spad_half_size; \
size_t dst_spad_half_size = actx->dst_spad_half_size; \
\
const int BLOCK = actx->block; \
if (BLOCK == 0) { \
FARF(ERROR, \
OP_STR \
" : current VTCM reservation %zu is too small for even 1 row per thread, needed at least %zu\n", \
actx->vtcm_src0_size_per_thread, src0_row_size_aligned); \
return; \
} \
\
dma_queue * dma_queue = actx->octx->ctx->dma[ith]; \
\
/* See discussion: https://github.com/ggml-org/llama.cpp/pull/18151#issuecomment-3678235379 */ \
for (uint32_t ir = src0_start_row, spad_idx = 0; ir < src0_end_row && spad_idx < 2; ir += BLOCK, spad_idx++) { \
const uint32_t block_size = MIN(BLOCK, src0_end_row - ir); \
\
/* Dummy DMA transation for sequencing (interleaving dst,src,dst,...) */ \
dma_queue_push_vtcm_to_ddr(dma_queue, \
dma_make_ptr(data_dst, dst_spad_data + (spad_idx * dst_spad_half_size)), \
dst_row_size, dst_row_size_aligned, 0); \
\
dma_queue_push( \
dma_queue, \
dma_make_ptr(src0_spad_data + (spad_idx * src0_spad_half_size), data_src0 + (ir * src0_row_stride)), \
src0_row_size_aligned, src0_row_stride, src0_row_size, block_size); \
dma_queue_push( \
dma_queue, \
dma_make_ptr(src1_spad_data + (spad_idx * src1_spad_half_size), data_src1 + (ir * src1_row_stride)), \
src1_row_size_aligned, src1_row_stride, src1_row_size, block_size); \
} \
\
for (uint32_t ir = src0_start_row; ir < src0_end_row; ir += BLOCK) { \
const uint32_t block_size = MIN(BLOCK, src0_end_row - ir); \
\
float * dst_spad = (float *) dma_queue_pop(dma_queue).src; \
float * src0_spad = (float *) dma_queue_pop(dma_queue).dst; \
float * src1_spad = (float *) dma_queue_pop(dma_queue).dst; \
\
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir); \
CORE_EXPR; \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, ir); \
\
dma_queue_push_vtcm_to_ddr(dma_queue, dma_make_ptr(data_dst + (ir * dst_row_size), dst_spad), \
dst_row_size, dst_row_size_aligned, block_size); \
\
/* prefetch N+2 loop iteration if any */ \
const uint32_t pref_block = (ir + BLOCK * 2); \
if (pref_block < src0_end_row) { \
const uint32_t pref_block_size = MIN(BLOCK, src0_end_row - pref_block); \
dma_queue_push(dma_queue, dma_make_ptr(src0_spad, data_src0 + (pref_block * src0_row_stride)), \
src0_row_size_aligned, src0_row_stride, src0_row_size, pref_block_size); \
dma_queue_push(dma_queue, dma_make_ptr(src1_spad, data_src1 + (pref_block * src1_row_stride)), \
src1_row_size_aligned, src1_row_stride, src1_row_size, pref_block_size); \
} \
} \
\
dma_queue_flush(dma_queue); \
\
}
DEFINE_GLU_PER_THREAD(swiglu, "swiglu-f32", swiglu_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
@@ -473,14 +473,30 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
}
const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads = MIN(octx->n_threads, src0_nrows);
const size_t dst_row_size = dst->ne[0] * SIZEOF_FP32;
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_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 = octx->n_threads;
// row_size = bytes of useful data per row (what the kernel touches / what DMA copies).
// row_stride = bytes between successive rows in DDR (may exceed row_size for non-contig src).
const size_t nc_bytes = dst->ne[0] * SIZEOF_FP32;
const size_t src0_row_size = nc_bytes;
const size_t src1_row_size = nc_bytes;
const size_t dst_row_size = nc_bytes;
const size_t nc_bytes = dst_row_size;
const size_t src0_row_size = nc_bytes;
const size_t src1_row_size = nc_bytes;
const size_t src0_row_stride = src0->nb[1];
const size_t src1_row_stride = src1 ? src1->nb[1] : src0->nb[1];
@@ -518,7 +534,7 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
struct htp_act_context actx;
actx.octx = octx;
actx.src0_nrows_per_thread = (src0_nrows + n_threads - 1) / n_threads;
actx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
actx.src0_row_size = src0_row_size;
actx.src1_row_size = src1_row_size;
@@ -545,7 +561,8 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
actx.dst_spad_half_size = L.dst_bytes_per_thread / 2;
actx.block = actx.src0_spad_half_size / actx.src0_row_size_aligned;
actx.src0_nrows = src0_nrows;
actx.src0_nrows = nrows;
actx.row_start = row_start;
actx.nc = dst->ne[0];
@@ -570,7 +587,7 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
actx.data_src1 = data_src1;
actx.data_dst = (uint8_t *) dst->data;
worker_pool_run_func(octx->ctx->worker_pool, act_op_func, &actx, n_threads);
work_queue_run(octx->ctx->work_queue, act_op_func, &actx, n_threads);
return HTP_STATUS_OK;
}
+86 -27
View File
@@ -17,6 +17,7 @@
#include "hex-dma.h"
#include "hex-profile.h"
#include "allreduce-ops.h"
#include "htp-fence.h"
struct htp_allreduce_context {
struct htp_ops_context * octx;
@@ -242,7 +243,42 @@ DEFINE_ALLREDUCE_THREAD_DMA_2D(add_f32, float, hvx_add_f32_aaa, 1, 0)
DEFINE_ALLREDUCE_THREAD_DMA_2D(add_bcast_f16, __fp16, hvx_add_f16_aaa, 1, 1)
DEFINE_ALLREDUCE_THREAD_DMA_2D(add_bcast_f32, float, hvx_add_f32_aaa, 1, 1)
static int validate_allreduce(
struct htp_ops_context * octx,
const struct htp_allreduce_kernel_params * kparams,
uint32_t n_ranks
) {
if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
if (kparams->vtcm_size_per_thread <= 0 || kparams->vtcm_size <= 0) {
return HTP_STATUS_INVAL_PARAMS;
}
const bool has_add = (octx->op == HTP_OP_ALLREDUCE_ADD);
const size_t n_vtcm_buffers = htp_allreduce_vtcm_buffer_count(
n_ranks, octx->n_threads, has_add, kparams->is_row_bcast != 0);
const size_t vtcm_size = n_vtcm_buffers * (size_t) kparams->vtcm_size_per_thread;
if (vtcm_size != (size_t) kparams->vtcm_size) {
return HTP_STATUS_INVAL_PARAMS;
}
if (vtcm_size > octx->ctx->vtcm_size) {
return HTP_STATUS_VTCM_TOO_SMALL;
}
if (octx->dst->type != HTP_TYPE_F16 && octx->dst->type != HTP_TYPE_F32) {
return HTP_STATUS_NO_SUPPORT;
}
return HTP_STATUS_OK;
}
int op_allreduce(struct htp_ops_context * octx) {
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
return HTP_STATUS_OK;
}
const struct htp_allreduce_kernel_params * kparams = (const struct htp_allreduce_kernel_params *) octx->kernel_params;
const struct htp_tensor * dst = octx->dst;
@@ -253,38 +289,53 @@ int op_allreduce(struct htp_ops_context * octx) {
return HTP_STATUS_INVAL_PARAMS;
}
if (dst->type != HTP_TYPE_F16 && dst->type != HTP_TYPE_F32) {
return HTP_STATUS_NO_SUPPORT;
}
const uint32_t nelem = dst->ne[0] * dst->ne[1] * dst->ne[2] * dst->ne[3];
const uint32_t fence_seq_entry = (uint32_t) octx->op_params[0];
const uint32_t fence_seq_exit = (uint32_t) octx->op_params[1];
const struct htp_tensor * my_sync = octx->src[n_ranks + rank];
atomic_uint * my_fence = (atomic_uint *) (uintptr_t) my_sync->data;
const int status = validate_allreduce(octx, kparams, n_ranks);
if (status != HTP_STATUS_OK) {
if (status == HTP_STATUS_NO_SUPPORT) {
FARF(ERROR, "ggml-hex: allreduce unsupported type %d : rank %u\n", dst->type, rank);
}
htp_fence_write(my_fence, fence_seq_exit, status);
return status;
}
const bool has_add = (octx->op == HTP_OP_ALLREDUCE_ADD);
const uint32_t nelem = dst->ne[0] * dst->ne[1] * dst->ne[2] * dst->ne[3];
// 1. Entry Barrier: Synchronize all ranks before reading
struct htp_thread_trace * tr0 = &octx->ctx->trace[0];
htp_trace_event_start(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry);
const struct htp_tensor * my_sync = octx->src[n_ranks + rank];
atomic_uint * my_fence = (atomic_uint *) my_sync->data;
atomic_store(&my_fence[0], fence_seq_entry);
asm volatile ("syncht" : : : "memory");
Q6_dccleaninva_A((void *) my_fence);
htp_fence_write(my_fence, fence_seq_entry, octx->status);
for (uint32_t j = 0; j < n_ranks; j++) {
if (j == rank) continue;
const struct htp_tensor * peer_sync = octx->src[n_ranks + j];
atomic_uint * peer_fence = (atomic_uint *) peer_sync->data;
atomic_uint * peer_fence = (atomic_uint *) (uintptr_t) peer_sync->data;
uint64_t spins = 0;
while (1) {
Q6_dccleaninva_A((void *) peer_fence);
uint32_t val = atomic_load(&peer_fence[0]);
if (val == fence_seq_entry || val == fence_seq_exit) {
uint32_t peer_seq;
uint32_t peer_status;
htp_fence_read(peer_fence, &peer_seq, &peer_status);
if ((int32_t)(peer_seq - fence_seq_entry) >= 0) {
if (peer_status > HTP_STATUS_OK) {
FARF(ERROR, "ggml-hex: allreduce entry peer %u failed with status %u\n", j, peer_status);
htp_fence_write(my_fence, fence_seq_exit, peer_status);
htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry);
return peer_status;
}
break;
}
if (++spins > HTP_FENCE_TIMEOUT) {
FARF(ERROR, "ggml-hex: allreduce entry fence-wait TIMEOUT: rank %u waiting on %u (fence %p seq %u)\n", rank, j, peer_fence, fence_seq_entry);
FARF(ERROR, "ggml-hex: allreduce entry fence-wait TIMEOUT : rank %u waiting on %u fence %p seq 0x%x peer-seq 0x%x\n",
rank, j, peer_fence, fence_seq_entry, peer_seq);
htp_fence_write(my_fence, fence_seq_exit, HTP_STATUS_INTERNAL_ERR);
htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_entry);
return HTP_STATUS_INTERNAL_ERR;
}
hex_pause();
@@ -301,8 +352,6 @@ int op_allreduce(struct htp_ops_context * octx) {
const uint32_t elems_per_thread = (uint32_t) kparams->elems_per_thread;
const uint32_t vtcm_size_per_thread = (uint32_t) kparams->vtcm_size_per_thread;
const bool has_add = (octx->op == HTP_OP_ALLREDUCE_ADD);
struct htp_allreduce_context actx;
actx.octx = octx;
actx.n_ranks = n_ranks;
@@ -339,6 +388,8 @@ int op_allreduce(struct htp_ops_context * octx) {
}
break;
default:
FARF(ERROR, "ggml-hex: allreduce unsupported kernel %d : rank %u\n", kparams->kernel_type, rank);
htp_fence_write(my_fence, fence_seq_exit, HTP_STATUS_NO_SUPPORT);
return HTP_STATUS_NO_SUPPORT;
}
@@ -368,23 +419,31 @@ int op_allreduce(struct htp_ops_context * octx) {
// 4. Exit Barrier: Synchronize all ranks after writing
htp_trace_event_start(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);
atomic_store(&my_fence[0], fence_seq_exit);
asm volatile ("syncht" : : : "memory");
Q6_dccleaninva_A((void *) my_fence);
htp_fence_write(my_fence, fence_seq_exit, octx->status);
for (uint32_t j = 0; j < n_ranks; j++) {
if (j == rank) continue;
const struct htp_tensor * peer_sync = octx->src[n_ranks + j];
atomic_uint * peer_fence = (atomic_uint *) peer_sync->data;
atomic_uint * peer_fence = (atomic_uint *) (uintptr_t) peer_sync->data;
uint64_t spins = 0;
while (1) {
Q6_dccleaninva_A((void *) peer_fence);
uint32_t val = atomic_load(&peer_fence[0]);
if (val == fence_seq_exit) {
uint32_t peer_seq;
uint32_t peer_status;
htp_fence_read(peer_fence, &peer_seq, &peer_status);
if ((int32_t)(peer_seq - fence_seq_exit) >= 0) {
if (peer_status > HTP_STATUS_OK) {
FARF(ERROR, "ggml-hex: allreduce exit peer %u failed with status %u\n", j, peer_status);
htp_fence_write(my_fence, fence_seq_exit, peer_status);
htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);
return peer_status;
}
break;
}
if (++spins > HTP_FENCE_TIMEOUT) {
FARF(ERROR, "ggml-hex: allreduce exit fence-wait TIMEOUT: rank %u waiting on %u (fence %p seq %u)\n", rank, j, peer_fence, fence_seq_exit);
FARF(ERROR, "ggml-hex: allreduce exit fence-wait TIMEOUT : rank %u waiting on %u fence %p seq 0x%x peer-seq 0x%x\n",
rank, j, peer_fence, fence_seq_exit, peer_seq);
htp_fence_write(my_fence, fence_seq_exit, HTP_STATUS_INTERNAL_ERR);
htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);
return HTP_STATUS_INTERNAL_ERR;
}
hex_pause();
@@ -394,5 +453,5 @@ int op_allreduce(struct htp_ops_context * octx) {
htp_trace_event_stop(tr0, HTP_TRACE_EVT_FENCE, (uint16_t) fence_seq_exit);
return HTP_STATUS_OK;
return octx->status;
}
+11
View File
@@ -2,6 +2,8 @@
#define ALLREDUCE_OPS_H
#include <stdint.h>
#include <stddef.h>
#include <stdbool.h>
#define HTP_ALLREDUCE_MAX_RANKS 4
@@ -15,6 +17,15 @@ enum htp_allreduce_kernel_type {
HTP_ALLREDUCE_KERNEL_DMA_2D,
};
static inline size_t htp_allreduce_vtcm_buffer_count(
uint32_t n_ranks,
uint32_t n_threads,
bool has_add,
bool is_row_bcast
) {
return (size_t) (n_ranks + 1) * n_threads + (has_add ? (is_row_bcast ? 1 : n_threads) : 0);
}
struct htp_allreduce_kernel_params {
int32_t rank;
int32_t n_ranks;
+36 -14
View File
@@ -11,9 +11,10 @@
#include "hvx-utils.h"
#include "hex-dma.h"
#include "hex-common.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#ifndef MIN
#define MIN(a, b) ((a) < (b) ? (a) : (b))
@@ -22,6 +23,9 @@
struct htp_argsort_context {
struct htp_ops_context * octx;
uint32_t nrows_per_thread;
uint32_t total_rows;
uint32_t row_start;
uint32_t row_end;
uint8_t * vtcm_base;
size_t vtcm_per_thread;
};
@@ -336,10 +340,9 @@ static void htp_argsort_f32_##ne00##_##order_name(unsigned int n, unsigned int i
const struct htp_tensor * src0 = octx->src[0]; \
const struct htp_tensor * dst = octx->dst; \
uint8_t * spad = actx->vtcm_base + actx->vtcm_per_thread * i; \
uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3]; \
uint32_t rows_per_thread = actx->nrows_per_thread; \
uint32_t start_row = rows_per_thread * i; \
uint32_t end_row = MIN(start_row + rows_per_thread, total_rows); \
uint32_t start_row = actx->row_start + rows_per_thread * i; \
uint32_t end_row = MIN(start_row + rows_per_thread, actx->row_end); \
size_t values_size = hex_round_up(ne00 * sizeof(float), 128); \
float * values_buf = (float *) spad; \
int32_t * indices_buf = (int32_t *) (spad + values_size); \
@@ -386,9 +389,6 @@ static void htp_argsort_f32_fallback(unsigned int n, unsigned int i, void * data
// Dimensions
uint32_t ne00 = src0->ne[0];
uint32_t ne01 = src0->ne[1];
uint32_t ne02 = src0->ne[2];
uint32_t ne03 = src0->ne[3];
uint32_t nb01 = src0->nb[1];
@@ -398,10 +398,9 @@ static void htp_argsort_f32_fallback(unsigned int n, unsigned int i, void * data
enum ggml_sort_order order = (enum ggml_sort_order) octx->op_params[0];
// Rows to process
uint32_t total_rows = ne01 * ne02 * ne03;
uint32_t rows_per_thread = actx->nrows_per_thread;
uint32_t start_row = rows_per_thread * i;
uint32_t end_row = MIN(start_row + rows_per_thread, total_rows);
uint32_t start_row = actx->row_start + rows_per_thread * i;
uint32_t end_row = MIN(start_row + rows_per_thread, actx->row_end);
size_t values_size = hex_round_up(ne00 * sizeof(float), 128);
uint32_t num_vec_ind_values = hmx_ceil_div(ne00, VLEN/(sizeof(int32_t)));
@@ -451,8 +450,28 @@ int op_argsort(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
const uint32_t total_rows = octx->src[0]->ne[1] * octx->src[0]->ne[2] * octx->src[0]->ne[3];
const uint32_t n_threads = MIN(total_rows, octx->n_threads);
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * dst = octx->dst;
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const size_t dst_row_size = dst->ne[0] * sizeof(int32_t);
uint32_t row_start = 0;
uint32_t row_end = total_rows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, sizeof(int32_t), (uint32_t) dst_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
row_start = range.start;
row_end = range.start + range.count;
}
const uint32_t nrows = row_end - row_start;
if (nrows == 0) {
return HTP_STATUS_OK;
}
const uint32_t n_threads = octx->n_threads;
// Allocate scratchpad
// We need 1 row of float + 1 row of int32 per thread.
@@ -478,7 +497,10 @@ int op_argsort(struct htp_ops_context * octx) {
struct htp_argsort_context actx;
actx.octx = octx;
actx.nrows_per_thread = (total_rows + n_threads - 1) / n_threads;
actx.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
actx.total_rows = nrows;
actx.row_start = row_start;
actx.row_end = row_end;
actx.vtcm_base = (uint8_t *) octx->ctx->vtcm_base;
actx.vtcm_per_thread = spad_per_thread;
@@ -508,7 +530,7 @@ int op_argsort(struct htp_ops_context * octx) {
}
// Run jobs
worker_pool_run_func(octx->ctx->worker_pool, job_func, &actx, n_threads);
work_queue_run(octx->ctx->work_queue, job_func, &actx, n_threads);
return HTP_STATUS_OK;
}
+136 -92
View File
@@ -13,9 +13,10 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#ifndef MIN
@@ -36,6 +37,8 @@ struct htp_binary_context {
uint32_t block_max;
uint32_t nrows_per_thread;
uint32_t total_rows;
uint32_t row_start;
size_t src0_row_size_aligned;
size_t src1_row_size_aligned;
size_t dst_row_size_aligned;
@@ -48,27 +51,27 @@ struct htp_binary_context {
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 ne00 = src0->ne[0]; \
const uint32_t ne01 = src0->ne[1]; \
const uint32_t ne02 = src0->ne[2]; \
const uint32_t ne03 = src0->ne[3]; \
\
const uint32_t ne10 = src1->ne[0]; \
const uint32_t ne11 = src1->ne[1]; \
const uint32_t ne12 = src1->ne[2]; \
const uint32_t ne13 = src1->ne[3]; \
\
const uint32_t nb01 = src0->nb[1]; \
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t nb11 = src1->nb[1]; \
const uint32_t nb12 = src1->nb[2]; \
const uint32_t nb13 = src1->nb[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
\
const uint32_t ne00 = src0->ne[0]; \
const uint32_t ne01 = src0->ne[1]; \
const uint32_t ne02 = src0->ne[2]; \
const uint32_t ne03 = src0->ne[3]; \
\
const uint32_t ne10 = src1->ne[0]; \
const uint32_t ne11 = src1->ne[1]; \
const uint32_t ne12 = src1->ne[2]; \
const uint32_t ne13 = src1->ne[3]; \
\
const uint32_t nb01 = src0->nb[1]; \
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t nb11 = src1->nb[1]; \
const uint32_t nb12 = src1->nb[2]; \
const uint32_t nb13 = src1->nb[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3];
static inline uint32_t calc_block_size(struct htp_binary_context * bctx, uint32_t ir, uint32_t end_row, uint32_t ne01, uint32_t ne02) {
@@ -93,87 +96,87 @@ static inline uint32_t calc_block_size(struct htp_binary_context * bctx, uint32_
}
// Macro for scalar op switch
#define COMPUTE_SCALAR_OP(DST, SRC, VAL, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
case HTP_OP_SUB: hvx_sub_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
case HTP_OP_MUL: hvx_mul_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
#define COMPUTE_SCALAR_OP(DST, SRC, VAL, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
case HTP_OP_SUB: hvx_sub_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
case HTP_OP_MUL: hvx_mul_scalar_f32_aa(DST, SRC, *(float *)VAL, N); break; \
case HTP_OP_DIV: hvx_mul_scalar_f32_aa(DST, SRC, 1.0f / (*(float *)VAL), N); break; \
default: break; \
} \
} \
else { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
case HTP_OP_SUB: hvx_sub_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
case HTP_OP_MUL: hvx_mul_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
case HTP_OP_DIV: hvx_div_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
default: break; \
} \
default: break; \
} \
} \
else { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
case HTP_OP_SUB: hvx_sub_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
case HTP_OP_MUL: hvx_mul_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
case HTP_OP_DIV: hvx_div_scalar_f16_aa(DST, SRC, *(_Float16 *)VAL, N); break; \
default: break; \
} \
}
// Macro for vector op switch (All Aligned)
#define COMPUTE_VECTOR_OP_AAA(DST, SRC0, SRC1, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
#define COMPUTE_VECTOR_OP_AAA(DST, SRC0, SRC1, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_f32_aaa(DST, SRC0, SRC1, N); break; \
case HTP_OP_SUB: hvx_sub_f32_aaa(DST, SRC0, SRC1, N); break; \
case HTP_OP_MUL: hvx_mul_f32_aaa(DST, SRC0, SRC1, N); break; \
case HTP_OP_DIV: hvx_div_f32_aaa(DST, SRC0, SRC1, N); break; \
default: break; \
} \
} \
else { \
switch (octx->op) { \
default: break; \
} \
} \
else { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_f16_aaa(DST, SRC0, SRC1, N); break; \
case HTP_OP_SUB: hvx_sub_f16_aaa(DST, SRC0, SRC1, N); break; \
case HTP_OP_MUL: hvx_mul_f16_aaa(DST, SRC0, SRC1, N); break; \
case HTP_OP_DIV: hvx_div_f16_aaa(DST, SRC0, SRC1, N); break; \
default: break; \
} \
default: break; \
} \
}
// Macro for vector op switch (Dst Aligned, Src0 Aligned, Src1 Unaligned)
#define COMPUTE_VECTOR_OP_AAU(DST, SRC0, SRC1, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
#define COMPUTE_VECTOR_OP_AAU(DST, SRC0, SRC1, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_f32_aau(DST, SRC0, SRC1, N); break; \
case HTP_OP_SUB: hvx_sub_f32_aau(DST, SRC0, SRC1, N); break; \
case HTP_OP_MUL: hvx_mul_f32_aau(DST, SRC0, SRC1, N); break; \
case HTP_OP_DIV: hvx_div_f32_aau(DST, SRC0, SRC1, N); break; \
default: break; \
} \
} \
else { \
switch (octx->op) { \
default: break; \
} \
} \
else { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_f16_aau(DST, SRC0, SRC1, N); break; \
case HTP_OP_SUB: hvx_sub_f16_aau(DST, SRC0, SRC1, N); break; \
case HTP_OP_MUL: hvx_mul_f16_aau(DST, SRC0, SRC1, N); break; \
case HTP_OP_DIV: hvx_div_f16_aau(DST, SRC0, SRC1, N); break; \
default: break; \
} \
default: break; \
} \
}
// Macro for vector op switch (All Unaligned - generic loop used in element repeat)
#define COMPUTE_VECTOR_OP_UUU(DST, SRC0, SRC1, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
#define COMPUTE_VECTOR_OP_UUU(DST, SRC0, SRC1, TYPE, N) \
if(TYPE == HTP_TYPE_F32) { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_f32_uuu(DST, SRC0, SRC1, N); break; \
case HTP_OP_SUB: hvx_sub_f32_uuu(DST, SRC0, SRC1, N); break; \
case HTP_OP_MUL: hvx_mul_f32_uuu(DST, SRC0, SRC1, N); break; \
case HTP_OP_DIV: hvx_div_f32_uuu(DST, SRC0, SRC1, N); break; \
default: break; \
} \
} \
else { \
switch (octx->op) { \
default: break; \
} \
} \
else { \
switch (octx->op) { \
case HTP_OP_ADD: hvx_add_f16_uuu(DST, SRC0, SRC1, N); break; \
case HTP_OP_SUB: hvx_sub_f16_uuu(DST, SRC0, SRC1, N); break; \
case HTP_OP_MUL: hvx_mul_f16_uuu(DST, SRC0, SRC1, N); break; \
case HTP_OP_DIV: hvx_div_f16_uuu(DST, SRC0, SRC1, N); break; \
default: break; \
} \
default: break; \
} \
}
// 1. Scalar src1 (ne10 == 1)
@@ -184,9 +187,8 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
const uint32_t src0_type = octx->src[0]->type;
const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
const uint32_t total_rows = ne01 * ne02 * ne03;
const uint32_t start_row = bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, total_rows);
const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
if (start_row >= end_row) return;
FARF(HIGH, "binary-scalar: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -222,6 +224,8 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
}
// Main loop
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = start_row; ir < end_row; ) {
uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
@@ -242,12 +246,14 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
uint8_t * src1_ptr = (uint8_t *)src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11;
uint32_t s1_stride = (ne11 == 1) ? 0 : nb11;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
for (uint32_t r = 0; r < current_block_size; r++) {
uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
uint8_t * r_dst = d_spad + r * bctx->dst_row_size_aligned;
COMPUTE_SCALAR_OP(r_dst, r_src0, src1_ptr, src0_type, ne00);
src1_ptr += s1_stride;
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
@@ -266,6 +272,7 @@ static void binary_job_scalar(unsigned int nth, unsigned int ith, void * data) {
}
ir += current_block_size;
}
dma_queue_flush(q);
}
@@ -277,9 +284,8 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi
const uint32_t src0_type = octx->src[0]->type;
const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
const uint32_t total_rows = ne01 * ne02 * ne03;
const uint32_t start_row = bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, total_rows);
const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
if (start_row >= end_row) return;
FARF(HIGH, "binary-same-shape: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -323,18 +329,22 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = start_row; ir < end_row; ) {
uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
uint8_t * s1_spad = (uint8_t *) dma_queue_pop(q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
for (uint32_t r = 0; r < current_block_size; r++) {
uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
uint8_t * r_src1 = s1_spad + r * bctx->src1_row_size_aligned;
uint8_t * r_dst = d_spad + r * bctx->dst_row_size_aligned;
COMPUTE_VECTOR_OP_AAA(r_dst, r_src0, r_src1, src0_type, ne00);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
uint32_t i03, i02, i01, rem;
i03 = fastdiv(ir, &bctx->src0_dim12_div);
@@ -366,6 +376,7 @@ static void binary_job_vector_same_shape(unsigned int nth, unsigned int ith, voi
}
ir += current_block_size;
}
dma_queue_flush(q);
}
@@ -377,9 +388,8 @@ static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith,
const uint32_t src0_type = octx->src[0]->type;
const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
const uint32_t total_rows = ne01 * ne02 * ne03;
const uint32_t start_row = bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, total_rows);
const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
if (start_row >= end_row) return;
FARF(HIGH, "binary-row-bcast: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -416,17 +426,21 @@ static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith,
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = start_row; ir < end_row; ) {
uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
uint8_t * s0_spad = (uint8_t *) dma_queue_pop(q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
for (uint32_t r = 0; r < current_block_size; r++) {
uint8_t * r_src0 = s0_spad + r * bctx->src0_row_size_aligned;
uint8_t * r_src1 = (uint8_t *)s1_ptr; // Constant
uint8_t * r_dst = d_spad + r * bctx->dst_row_size_aligned;
COMPUTE_VECTOR_OP_AAA(r_dst, r_src0, r_src1, src0_type, ne00);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
uint32_t i03 = fastdiv(ir, &bctx->src0_dim12_div);
uint32_t rem = ir - i03 * (ne02 * ne01);
@@ -447,6 +461,7 @@ static void binary_job_vector_row_broadcast(unsigned int nth, unsigned int ith,
}
ir += current_block_size;
}
dma_queue_flush(q);
}
@@ -458,9 +473,8 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
const uint32_t src0_type = octx->src[0]->type;
const uint32_t row_size_bytes = (src0_type == HTP_TYPE_F32) ? ne00 * sizeof(float) : ne00 * sizeof(_Float16);
const uint32_t total_rows = ne01 * ne02 * ne03;
const uint32_t start_row = bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, total_rows);
const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
if (start_row >= end_row) return;
FARF(HIGH, "binary-complex: %d/%d (%u:%u) row-size %u (%u)", ith, nth, start_row, end_row, nb01, bctx->dst_row_size_aligned);
@@ -493,6 +507,8 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = start_row; ir < end_row; ) {
uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
@@ -503,6 +519,7 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
uint32_t i01 = rem - i02 * ne01;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
for (uint32_t r = 0; r < current_block_size; r++) {
uint32_t r_i01 = i01 + r;
uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
@@ -516,6 +533,7 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
// Read src1 from DDR (unaligned)
COMPUTE_VECTOR_OP_AAU(r_dst, r_src0, r_src1, src0_type, ne00);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
@@ -532,6 +550,7 @@ static void binary_job_vector_complex(unsigned int nth, unsigned int ith, void *
}
ir += current_block_size;
}
dma_queue_flush(q);
}
@@ -544,9 +563,8 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
const uint32_t src0_type = octx->src[0]->type;
const uint32_t elem_size_bytes = (src0_type == HTP_TYPE_F32) ? sizeof(float) : sizeof(_Float16);
const uint32_t row_size_bytes = ne00 * elem_size_bytes;;
const uint32_t total_rows = ne01 * ne02 * ne03;
const uint32_t start_row = bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, total_rows);
const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
if (start_row >= end_row) return;
uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
@@ -579,6 +597,8 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = start_row; ir < end_row; ) {
uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
@@ -589,6 +609,7 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
uint32_t i01 = rem - i02 * ne01;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
for (uint32_t r = 0; r < current_block_size; r++) {
uint32_t r_i01 = i01 + r;
uint32_t i13 = fastmodulo(i03, ne13, &bctx->src1_dim3_div);
@@ -606,6 +627,7 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
COMPUTE_VECTOR_OP_UUU(r_dst + c * elem_size_bytes, r_src0 + c * elem_size_bytes, r_src1_row, src0_type, len);
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, row_size_bytes, current_block_size);
@@ -622,6 +644,7 @@ static void binary_job_element_repeat(unsigned int nth, unsigned int ith, void *
}
ir += current_block_size;
}
dma_queue_flush(q);
}
@@ -650,9 +673,8 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
const uint32_t nb2 = dst->nb[2];
const uint32_t nb3 = dst->nb[3];
const uint32_t total_rows = ne01 * ne02 * ne03;
const uint32_t start_row = bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, total_rows);
const uint32_t start_row = bctx->row_start + bctx->nrows_per_thread * ith;
const uint32_t end_row = MIN(start_row + bctx->nrows_per_thread, bctx->row_start + bctx->total_rows);
if (start_row >= end_row) return;
uint8_t * src0_spad_base = octx->src0_spad.data + (ith * octx->src0_spad.size_per_thread);
@@ -683,6 +705,8 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = start_row; ir < end_row; ) {
uint32_t current_block_size = calc_block_size(bctx, ir, end_row, ne01, ne02);
uint8_t * d_spad = (uint8_t *) dma_queue_pop(q).src;
@@ -693,6 +717,7 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
uint32_t i02 = fastdiv(rem, &bctx->src0_dim1_div);
uint32_t i01 = rem - i02 * ne01;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
for (uint32_t r = 0; r < current_block_size; r++) {
uint32_t r_i01 = i01 + r; // linear within block since we split at ne01
@@ -704,6 +729,7 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
hvx_add_f32_aau(r_dst, r_src0, r_src1, ne00);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
uint8_t * dst_curr = (uint8_t *)dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1;
dma_queue_push(q, dma_make_ptr(dst_curr, d_spad), nb1, bctx->dst_row_size_aligned, ne00 * sizeof(float), current_block_size);
@@ -720,6 +746,7 @@ static void binary_job_add_id(unsigned int nth, unsigned int ith, void * data) {
}
ir += current_block_size;
}
dma_queue_flush(q);
}
@@ -729,15 +756,31 @@ static int execute_op_binary(struct htp_ops_context * octx) {
const struct htp_tensor * dst = octx->dst;
const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads = MIN(octx->n_threads, src0_nrows);
// Use packed row sizes for VTCM allocation
// Use packed row sizes for VTCM allocation and alignment
const uint32_t src0_type = octx->src[0]->type;
const size_t elem_size = (src0_type == HTP_TYPE_F32) ? sizeof(float) : sizeof(_Float16);
const size_t src0_row_size = src0->ne[0] * elem_size;
const size_t src1_row_size = src1->ne[0] * elem_size;
const size_t dst_row_size = dst->ne[0] * elem_size;
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_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 = octx->n_threads;
size_t src0_row_size_aligned = hex_round_up(src0_row_size, VLEN);
size_t src1_row_size_aligned = hex_round_up(src1_row_size, VLEN);
size_t dst_row_size_aligned = hex_round_up(dst_row_size, VLEN);
@@ -815,7 +858,9 @@ static int execute_op_binary(struct htp_ops_context * octx) {
struct htp_binary_context bctx;
bctx.octx = octx;
bctx.nrows_per_thread = (src0_nrows + n_threads - 1) / n_threads;
bctx.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
bctx.total_rows = nrows;
bctx.row_start = row_start;
bctx.block_max = rows_per_buffer;
bctx.src0_row_size_aligned = src0_row_size_aligned;
bctx.src1_row_size_aligned = src1_row_size_aligned;
@@ -850,7 +895,7 @@ static int execute_op_binary(struct htp_ops_context * octx) {
dma_queue_pop(q);
}
worker_pool_run_func(octx->ctx->worker_pool, worker_func, &bctx, n_threads);
work_queue_run(octx->ctx->work_queue, worker_func, &bctx, n_threads);
return HTP_STATUS_OK;
}
@@ -870,4 +915,3 @@ int op_binary(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
+66 -15
View File
@@ -1,5 +1,8 @@
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#include "hexagon_types.h"
#include "hexagon_protos.h"
#include "hvx_hexagon_protos.h"
@@ -13,6 +16,10 @@ struct htp_concat_context {
struct htp_ops_context * octx;
uint32_t dim;
uint32_t nrows_per_thread;
uint32_t row_start;
uint32_t nrows;
uint32_t elem_start;
uint32_t nelems;
struct fastdiv_values div_ne0;
struct fastdiv_values div_ne1;
struct fastdiv_values div_ne2;
@@ -28,10 +35,10 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
const uint32_t src0_ne0 = src0->ne[0];
const uint32_t src1_ne0 = src1->ne[0];
const uint32_t ne1 = dst->ne[1];
const uint32_t start_i = ith * cctx->nrows_per_thread;
const uint32_t end_i = (start_i + cctx->nrows_per_thread < ne1) ? (start_i + cctx->nrows_per_thread) : ne1;
const uint32_t row_end = cctx->row_start + cctx->nrows;
const uint32_t start_i = cctx->row_start + ith * cctx->nrows_per_thread;
const uint32_t end_i = (start_i + cctx->nrows_per_thread < row_end) ? (start_i + cctx->nrows_per_thread) : row_end;
if (start_i >= end_i) return;
dma_queue * q = octx->ctx->dma[ith];
@@ -51,6 +58,8 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
const uint32_t spad0_row_bytes = hex_round_up((src0_ne0 + src1_ne0_padded) * sizeof(float), VLEN);
uint32_t mu = src1_ne0_padded * spad1_stride;
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;
@@ -66,6 +75,7 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
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++) {
@@ -75,6 +85,7 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
dma_queue_pop(q); // src0
@@ -95,10 +106,10 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
const uint32_t src0_ne0 = src0->ne[0];
const uint32_t src1_ne0 = src1->ne[0];
const uint32_t ne1 = dst->ne[1];
const uint32_t start_i = ith * cctx->nrows_per_thread;
const uint32_t end_i = (start_i + cctx->nrows_per_thread < ne1) ? (start_i + cctx->nrows_per_thread) : ne1;
const uint32_t row_end = cctx->row_start + cctx->nrows;
const uint32_t start_i = cctx->row_start + ith * cctx->nrows_per_thread;
const uint32_t end_i = (start_i + cctx->nrows_per_thread < row_end) ? (start_i + cctx->nrows_per_thread) : row_end;
if (start_i >= end_i) return;
dma_queue * q = octx->ctx->dma[ith];
@@ -118,6 +129,8 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
const uint32_t spad0_row_bytes = hex_round_up((src0_ne0 + src1_ne0_padded) * sizeof(__fp16), VLEN);
uint32_t mu = src1_ne0_padded * spad1_stride;
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;
@@ -133,6 +146,7 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
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++) {
@@ -142,6 +156,7 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
dma_queue_pop(q); // src0
@@ -164,11 +179,14 @@ static void concat_generic(unsigned int nth, unsigned int ith, void * data) {
const uint32_t type_size = (dst->type == HTP_TYPE_F32 || dst->type == HTP_TYPE_I32) ? 4 : 2;
const uint32_t ne[4] = {dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]};
const uint32_t total_elements = ne[0] * ne[1] * ne[2] * ne[3];
const uint32_t chunk_size = (total_elements + nth - 1) / nth;
const uint32_t start_idx = MIN(ith * chunk_size, total_elements);
const uint32_t end_idx = MIN(start_idx + chunk_size, total_elements);
// Per-device element range aligned to prevent false sharing
const uint32_t elem_start = cctx->elem_start;
const uint32_t nelems = cctx->nelems;
const uint32_t chunk_size = (nelems + nth - 1) / nth;
const uint32_t start_idx = MIN(elem_start + ith * chunk_size, elem_start + nelems);
const uint32_t end_idx = MIN(start_idx + chunk_size, elem_start + nelems);
// Naive scalar element-wise copy
for (uint32_t idx = start_idx; idx < end_idx; idx++) {
@@ -236,13 +254,28 @@ 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) {
n_threads = MIN(dst->ne[1], n_threads);
if (n_threads < 1) {
n_threads = 1;
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;
uint32_t nrows = total_rows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, type_size, (uint32_t) dst_data_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_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;
}
cctx.row_start = row_start;
cctx.nrows = nrows;
uint32_t block_i = (type_size == 4) ? 32 : 64;
cctx.nrows_per_thread = hmx_ceil_div(dst->ne[1], n_threads);
cctx.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
// Allocate VTCM
uint32_t spad1_stride = block_i * type_size;
@@ -270,8 +303,26 @@ int op_concat(struct htp_ops_context * octx) {
} else {
worker_func = concat_2d_f16_transposed;
}
} else {
const uint32_t total_elements = dst->ne[0] * dst->ne[1] * dst->ne[2] * dst->ne[3];
uint32_t elem_start = 0;
uint32_t nelems = total_elements;
if (octx->ctx->mdev.count > 1) {
const uint32_t elems_per_chunk = HEX_L2_LINE_SIZE / type_size;
const bool can_split = htp_tensor_mdev_data_aligned(dst) && htp_tensor_is_contiguous(dst, type_size) && !htp_tensor_is_permuted(dst);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_elements, can_split ? elems_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
elem_start = range.start;
nelems = range.count;
}
if (nelems == 0) {
return HTP_STATUS_OK;
}
cctx.elem_start = elem_start;
cctx.nelems = nelems;
}
worker_pool_run_func(octx->ctx->worker_pool, worker_func, &cctx, n_threads);
work_queue_run(octx->ctx->work_queue, worker_func, &cctx, n_threads);
return HTP_STATUS_OK;
}
+279 -197
View File
@@ -16,6 +16,7 @@
#include "htp-ops.h"
#include "hvx-utils.h"
#include "htp-tensor.h"
#include "htp-fence.h"
struct htp_copy_context {
struct htp_ops_context * octx;
@@ -29,7 +30,23 @@ struct htp_copy_context {
uint32_t src0_blocks_per_row;
uint32_t dst_blocks_per_row;
uint32_t elem_start;
uint32_t nelem;
uint32_t elem_per_thread;
uint32_t src0_nrows_per_thread;
uint32_t row_start;
uint32_t nrows;
struct fastdiv_values div_ne01;
struct fastdiv_values div_ne02_ne01;
struct fastdiv_values div_ne0;
struct fastdiv_values div_ne1_ne0;
struct fastdiv_values div_ne2_ne1_ne0;
struct fastdiv_values div_ne00;
struct fastdiv_values div_ne01_ne00;
struct fastdiv_values div_ne02_ne01_ne00;
};
#define cpy_preamble \
@@ -54,131 +71,113 @@ struct htp_copy_context {
const uint32_t nb0 = dst->nb[0]; \
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3]; \
\
const uint32_t nr = ne01;
const uint32_t nb3 = dst->nb[3];
#define DEFINE_CPY_SAMESHAPE(NAME, ELEM_TYPE, ELEM_SIZE) \
static void cpy_thread_##NAME##_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 = dr * ith; \
const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr; \
if (ir0 >= nr) return; \
for (uint32_t i03 = 0; i03 < ne03; i03++) { \
for (uint32_t i02 = 0; i02 < ne02; i02++) { \
_Pragma("unroll(4)") \
for (uint32_t i01 = ir0; i01 < ir1; i01++) { \
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; \
hex_l2fetch(src0_ptr, ne00 * ELEM_SIZE, nb01, 2); \
hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE); \
} \
} \
} \
#define DEFINE_CPY_SAMESHAPE(NAME, ELEM_TYPE, ELEM_SIZE) \
static void cpy_thread_##NAME##_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 bool contiguous = (nb01 == ne00 * ELEM_SIZE) && (nb1 == nb01) && \
(nb02 == ne01 * nb01) && (nb2 == nb02) && \
(nb03 == ne02 * nb02) && (nb3 == nb03); \
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; \
if (contiguous) { \
hvx_copy_uu(dst_ptr, src0_ptr, (ir1 - ir0) * ne00, ELEM_SIZE); \
return; \
} \
for (uint32_t r = ir0; r < ir1; r++) { \
hex_l2fetch(src0_ptr, ne00 * ELEM_SIZE, nb01, 2); \
hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE); \
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; \
} \
} \
}
DEFINE_CPY_SAMESHAPE(f32, float, 4)
DEFINE_CPY_SAMESHAPE(f16, __fp16, 2)
#define DEFINE_CPY_RESHAPE(NAME, ELEM_TYPE, ELEM_SIZE) \
static void cpy_thread_##NAME##_reshape(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 = dr * ith; \
const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr; \
if (ir0 >= nr) return; \
const bool src0_contig = (nb00 == ELEM_SIZE) && \
(nb01 == ne00 * nb00) && \
(nb02 == ne01 * nb01) && \
(nb03 == ne02 * nb02); \
const bool dst_contig = (nb0 == ELEM_SIZE) && \
(nb1 == ne0 * nb0) && \
(nb2 == ne1 * nb1) && \
(nb3 == ne2 * nb2); \
if (src0_contig && dst_contig) { \
for (int64_t i03 = 0; i03 < ne03; i03++) { \
for (int64_t i02 = 0; i02 < ne02; i02++) { \
uint8_t * src_ptr = (uint8_t *) src0->data + i03*nb03 + i02*nb02 + ir0*nb01; \
uint32_t flat = ((i03*ne02 + i02)*ne01 + ir0) * ne00; \
uint8_t * dst_ptr = (uint8_t *) dst->data + flat * ELEM_SIZE; \
hvx_copy_uu(dst_ptr, src_ptr, (ir1 - ir0) * ne00, ELEM_SIZE); \
} \
} \
return; \
} \
const bool reshape_flat_fast = (ne03 == 1 && ne2 == 1 && ne3 == 1) && \
(ne0 == ne00 * ne01) && (ne1 == ne02) && \
(nb00 == ELEM_SIZE) && (nb0 == ELEM_SIZE); \
if (reshape_flat_fast) { \
for (uint32_t i02 = 0; i02 < ne02; i02++) { \
for (uint32_t i01 = ir0; i01 < ir1; i01++) { \
uint8_t * src0_ptr = (uint8_t *) src0->data + i01 * nb01 + i02 * nb02; \
uint8_t * dst_ptr = (uint8_t *) dst->data + i01 * ne00 * ELEM_SIZE + i02 * nb1; \
hvx_copy_uu(dst_ptr, src0_ptr, ne00, ELEM_SIZE); \
} \
} \
return; \
} \
int64_t k10 = 0; \
int64_t i11 = 0; \
int64_t i12 = 0; \
int64_t i13 = 0; \
const int64_t nk00 = ct->src0_blocks_per_row; \
const int64_t nk0 = ct->dst_blocks_per_row; \
for (int64_t i03 = 0; i03 < ne03; i03++) { \
for (int64_t i02 = 0; i02 < ne02; i02++) { \
k10 += nk00 * ir0; \
while (k10 >= nk0) { \
k10 -= nk0; \
if (++i11 == ne1) { \
i11 = 0; \
if (++i12 == ne2) { \
i12 = 0; \
if (++i13 == ne3) { \
i13 = 0; \
} \
} \
} \
} \
for (int64_t i01 = ir0; i01 < ir1; i01++) { \
for (int64_t k00 = 0; k00 < nk00; k00++) { \
const char * src0_ptr = ((char *) src0->data + k00*nb00 + i01*nb01 + i02*nb02 + i03*nb03); \
char * dst_ptr = ((char *) dst->data + k10*nb0 + i11*nb1 + i12*nb2 + i13*nb3); \
memcpy(dst_ptr, src0_ptr, ELEM_SIZE); \
if (++k10 == nk0) { \
k10 = 0; \
if (++i11 == ne1) { \
i11 = 0; \
if (++i12 == ne2) { \
i12 = 0; \
if (++i13 == ne3) { \
i13 = 0; \
} \
} \
} \
} \
} \
} \
k10 += nk00 * (ne01 - ir1); \
while (k10 >= nk0) { \
k10 -= nk0; \
if (++i11 == ne1) { \
i11 = 0; \
if (++i12 == ne2) { \
i12 = 0; \
if (++i13 == ne3) { \
i13 = 0; \
} \
} \
} \
} \
} \
} \
#define DEFINE_CPY_RESHAPE(NAME, ELEM_TYPE, ELEM_SIZE) \
static void cpy_thread_##NAME##_reshape(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 th_nelem = ct->elem_per_thread; \
const uint32_t th_start = ct->elem_start + ith * th_nelem; \
const uint32_t th_end = MIN(th_start + th_nelem, ct->elem_start + ct->nelem); \
if (th_start >= th_end) return; \
\
const uint32_t ne01_ne00 = ne01 * ne00; \
const uint32_t ne02_ne01_ne00 = ne02 * ne01_ne00; \
const uint32_t ne1_ne0 = ne1 * ne0; \
const uint32_t ne2_ne1_ne0 = ne2 * ne1_ne0; \
\
uint32_t e = th_start; \
uint32_t i13 = fastdiv(e, &ct->div_ne2_ne1_ne0); \
uint32_t rem = e - i13 * ne2_ne1_ne0; \
uint32_t i12 = fastdiv(rem, &ct->div_ne1_ne0); \
uint32_t rem2 = rem - i12 * ne1_ne0; \
uint32_t i11 = fastdiv(rem2, &ct->div_ne0); \
uint32_t i10 = rem2 - i11 * ne0; \
\
uint32_t i03 = fastdiv(e, &ct->div_ne02_ne01_ne00); \
uint32_t rem_s = e - i03 * ne02_ne01_ne00; \
uint32_t i02 = fastdiv(rem_s, &ct->div_ne01_ne00); \
uint32_t rem2_s = rem_s - i02 * ne01_ne00; \
uint32_t i01 = fastdiv(rem2_s, &ct->div_ne00); \
uint32_t i00 = rem2_s - i01 * ne00; \
\
char * dst_ptr = (char *) dst->data + i10*nb0 + i11*nb1 + i12*nb2 + i13*nb3; \
const char * src0_ptr = (const char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03; \
\
for (; e < th_end; e++) { \
*((ELEM_TYPE *) dst_ptr) = *((const ELEM_TYPE *) src0_ptr); \
\
dst_ptr += nb0; \
if (++i10 == ne0) { \
i10 = 0; \
if (++i11 == ne1) { \
i11 = 0; \
if (++i12 == ne2) { \
i12 = 0; \
i13++; \
} \
} \
dst_ptr = (char *) dst->data + i11*nb1 + i12*nb2 + i13*nb3; \
} \
\
src0_ptr += nb00; \
if (++i00 == ne00) { \
i00 = 0; \
if (++i01 == ne01) { \
i01 = 0; \
if (++i02 == ne02) { \
i02 = 0; \
i03++; \
} \
} \
src0_ptr = (const char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03; \
} \
} \
}
DEFINE_CPY_RESHAPE(f32, float, 4)
@@ -189,22 +188,33 @@ static void cpy_thread_f16_f32_sameshape(unsigned int nth, unsigned int ith, voi
struct htp_ops_context * octx = ct->octx;
cpy_preamble;
// parallelize by src0 rows
const uint32_t dr = ct->src0_nrows_per_thread;
const uint32_t ir0 = dr * ith;
const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr;
if (ir0 >= nr) return;
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;
// copy by rows
for (uint32_t i03 = 0; i03 < ne03; i03++) {
for (uint32_t i02 = 0; i02 < ne02; i02++) {
#pragma unroll(2)
for (uint32_t i01 = ir0; i01 < ir1; i01++) {
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;
hex_l2fetch(src0_ptr, ne00 * sizeof(float), nb01, 2);
hvx_copy_f16_f32_uu(dst_ptr, src0_ptr, ne00);
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);
hvx_copy_f16_f32_uu(dst_ptr, src0_ptr, ne00);
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;
}
}
}
@@ -214,22 +224,33 @@ static void cpy_thread_f32_f16_sameshape(unsigned int nth, unsigned int ith, voi
struct htp_ops_context * octx = ct->octx;
cpy_preamble;
// parallelize by src0 rows
const uint32_t dr = ct->src0_nrows_per_thread;
const uint32_t ir0 = dr * ith;
const uint32_t ir1 = (ir0 + dr) < nr ? (ir0 + dr) : nr;
if (ir0 >= nr) return;
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;
// copy by rows
for (uint32_t i03 = 0; i03 < ne03; i03++) {
for (uint32_t i02 = 0; i02 < ne02; i02++) {
#pragma unroll(2)
for (uint32_t i01 = ir0; i01 < ir1; i01++) {
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;
hex_l2fetch(src0_ptr, ne00 * sizeof(__fp16), nb01, 2);
hvx_copy_f32_f16_uu(dst_ptr, src0_ptr, ne00);
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(__fp16), nb01, 2);
hvx_copy_f32_f16_uu(dst_ptr, src0_ptr, ne00);
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;
}
}
}
@@ -250,15 +271,19 @@ static inline void cpy_dma_sametype_sameshape(
dma_queue * q = octx->ctx->dma[0];
if (contiguous_outer) {
dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03);
dma_queue_pop(q);
if (!dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03)) {
dma_queue_flush(q);
dma_queue_push(q, dma_make_ptr((void *) dst->data, (const void *) src0->data), nb1, nb01, ne00 * elem_size, ne01 * ne02 * ne03);
}
dma_queue_flush(q);
return;
}
for (uint32_t i03 = 0; i03 < ne03; i03++) {
for (uint32_t i02 = 0; i02 < ne02; i02++) {
uint8_t* dst_ptr = (uint8_t*) dst->data + i02*nb2 + i03*nb3;
uint8_t* src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
uint8_t * dst_ptr = (uint8_t *) dst->data + i02 * nb2 + i03 * nb3;
uint8_t * src0_ptr = (uint8_t *) src0->data + i02 * nb02 + i03 * nb03;
if (!dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01)) {
dma_queue_flush(q);
dma_queue_push(q, dma_make_ptr(dst_ptr, src0_ptr), nb1, nb01, ne00 * elem_size, ne01);
@@ -269,10 +294,9 @@ static inline void cpy_dma_sametype_sameshape(
dma_queue_flush(q);
}
int op_cpy(struct htp_ops_context * octx) {
static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
cpy_preamble;
const uint32_t n_threads = MIN(nr, octx->n_threads);
*use_dma = false;
struct htp_copy_context ct;
ct.octx = octx;
@@ -296,59 +320,117 @@ int op_cpy(struct htp_ops_context * octx) {
}
const bool sametype = (src0->type == dst->type);
const bool transposed = (nb00 > nb01) || (nb0 > nb1);
const bool transposed = (nb00 > nb01) || (nb0 > nb1) ||
(nb00 != ct.src0_type_size) || (nb0 != ct.dst_type_size) ||
(nb01 < ne00 * ct.src0_type_size) || (nb1 < ne0 * ct.dst_type_size);
const bool sameshape = !transposed && (ne00 == ne0 && ne01 == ne1 && ne02 == ne2 && ne03 == ne3);
ct.src0_nrows_per_thread = (nr + n_threads - 1) / n_threads;
const uint32_t n_threads = octx->n_threads;
worker_callback_t copy_fun = NULL;
bool use_dma = false;
const bool dst_is_contiguous = htp_tensor_is_contiguous(dst, ct.dst_type_size);
if (sametype && sameshape) {
use_dma = true;
} else if (sameshape) {
/**/ if (dst->type == HTP_TYPE_F16 && src0->type == HTP_TYPE_F32)
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
return HTP_STATUS_NO_SUPPORT;
} else if (sametype) {
if (src0->type == HTP_TYPE_F32) {
copy_fun = cpy_thread_f32_reshape;
} else {
copy_fun = cpy_thread_f16_reshape;
if (sameshape) {
const uint32_t total_rows = ne01 * ne02 * ne03;
const uint32_t row_size = ne00 * ct.dst_type_size;
ct.div_ne01 = init_fastdiv_values(ne01);
ct.div_ne02_ne01 = init_fastdiv_values(ne02 * ne01);
uint32_t row_start = 0;
uint32_t nrows = total_rows;
if (octx->ctx->mdev.count > 1) {
const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
const bool can_split = htp_tensor_mdev_data_aligned(dst) && dst_is_contiguous;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, can_split ? rows_per_chunk : 0, 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;
}
ct.row_start = row_start;
ct.nrows = nrows;
ct.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
if (sametype && octx->ctx->mdev.count <= 1) {
*use_dma = true;
cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3);
} else {
work_queue_func_t copy_fun = NULL;
if (sametype) {
copy_fun = (src0->type == HTP_TYPE_F32) ? cpy_thread_f32_sameshape : cpy_thread_f16_sameshape;
} else if (dst->type == HTP_TYPE_F16 && src0->type == HTP_TYPE_F32) {
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 {
return HTP_STATUS_NO_SUPPORT;
}
work_queue_run(octx->ctx->work_queue, copy_fun, &ct, n_threads);
}
} else if (sametype) {
const uint32_t total_elems = ne0 * ne1 * ne2 * ne3;
const uint32_t elems_per_line = (ct.dst_type_size == 4) ? 32 : 64;
ct.div_ne0 = init_fastdiv_values(ne0);
ct.div_ne1_ne0 = init_fastdiv_values(ne1 * ne0);
ct.div_ne2_ne1_ne0 = init_fastdiv_values(ne2 * ne1 * ne0);
ct.div_ne00 = init_fastdiv_values(ne00);
ct.div_ne01_ne00 = init_fastdiv_values(ne01 * ne00);
ct.div_ne02_ne01_ne00 = init_fastdiv_values(ne02 * ne01 * ne00);
uint32_t elem_start = 0;
uint32_t nelem = total_elems;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) && dst_is_contiguous;
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;
}
ct.elem_start = elem_start;
ct.nelem = nelem;
ct.elem_per_thread = fastdiv(nelem + n_threads - 1, &octx->n_threads_div);
work_queue_func_t copy_fun = (src0->type == HTP_TYPE_F32) ? cpy_thread_f32_reshape : cpy_thread_f16_reshape;
work_queue_run(octx->ctx->work_queue, copy_fun, &ct, n_threads);
} else {
return HTP_STATUS_NO_SUPPORT;
}
FARF(HIGH, "cpy-%s-%s: (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_dma=%d n_threads %u\n",
src0->type == HTP_TYPE_F32 ? "f32" : "f16", dst->type == HTP_TYPE_F32 ? "f32" : "f16",
ne00, ne01, ne02, ne03, ne0, ne1, ne2, ne3, use_dma, n_threads);
if (use_dma) {
cpy_dma_sametype_sameshape(octx, dst, src0, ct.src0_type_size, ne00, ne01, ne02, ne03, nb01, nb02, nb03, nb1, nb2, nb3);
} else {
worker_pool_run_func(octx->ctx->worker_pool, copy_fun, &ct, n_threads);
}
const struct htp_tensor *sync = octx->src[1];
if (sync && (sync->flags & HTP_TENSOR_FENCE)) {
if (!use_dma) {
// htp_tensor_flush_all(octx->ctx, octx->dsts, 1);
qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE);
}
atomic_uint * sync_fence = (atomic_uint *) sync->data;
const uint32_t seq = (uint32_t) octx->op_params[0];
atomic_store(&sync_fence[0], seq);
asm volatile ("syncht" : : : "memory");
Q6_dccleaninva_A((void *) sync_fence);
FARF(HIGH, "ggml-hex: sync-release : fence %p seq %u\n", sync_fence, seq);
}
return HTP_STATUS_OK;
}
int op_cpy(struct htp_ops_context * octx) {
bool use_dma = false;
int status = exec_cpy(octx, &use_dma);
htp_ops_context_set_status(octx, status);
if (octx->op == HTP_OP_CPY_FENCE) {
if (!use_dma) {
htp_flush_dirty_ranges(octx->ctx);
}
htp_mdev_group_barrier(octx);
if (octx->ctx->mdev.idx == 0) {
const struct htp_tensor * sync = octx->src[1];
const uint32_t seq = (uint32_t) octx->op_params[0];
atomic_uint * sync_fence = (atomic_uint *) (uintptr_t) sync->data;
htp_fence_write(sync_fence, seq, octx->status);
FARF(HIGH, "ggml-hex: sync-release : fence %p seq 0x%x status %d\n", sync_fence, seq, octx->status);
}
}
return octx->status;
}
+62 -43
View File
@@ -7,6 +7,8 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
@@ -17,25 +19,25 @@
#define htp_cumsum_tensors_preamble \
const struct htp_tensor * restrict src0 = octx->src[0]; \
const struct htp_tensor * restrict dst = octx->dst; \
\
const uint32_t ne00 = src0->ne[0]; \
const uint32_t ne01 = src0->ne[1]; \
const uint32_t ne02 = src0->ne[2]; \
const uint32_t ne03 = src0->ne[3]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb00 = src0->nb[0]; \
const uint32_t nb01 = src0->nb[1]; \
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t nb0 = dst->nb[0]; \
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
\
const uint32_t ne00 = src0->ne[0]; \
const uint32_t ne01 = src0->ne[1]; \
const uint32_t ne02 = src0->ne[2]; \
const uint32_t ne03 = src0->ne[3]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb00 = src0->nb[0]; \
const uint32_t nb01 = src0->nb[1]; \
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t nb0 = dst->nb[0]; \
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3];
struct htp_cumsum_context {
@@ -46,6 +48,7 @@ struct htp_cumsum_context {
size_t dst_row_size_aligned;
uint32_t rows_per_thread;
uint32_t total_rows;
uint32_t row_start;
};
#define htp_cumsum_preamble \
@@ -116,11 +119,8 @@ static inline void hvx_cumsum_row_f32(const float * restrict src, float * restri
static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * data) {
htp_cumsum_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
const uint32_t ir0 = cctx->rows_per_thread * ith;
const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->total_rows);
const uint32_t ir0 = cctx->row_start + cctx->rows_per_thread * ith;
const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->row_start + cctx->total_rows);
if (ir0 >= ir1) {
return;
@@ -149,11 +149,15 @@ static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * dat
src_row_size_aligned, src_row_size, 1);
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = ir0; ir < ir1; ir++) {
float * dst_spad_row = (float *) dma_queue_pop(dma_queue).src;
float * src_spad_row = (float *) dma_queue_pop(dma_queue).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
hvx_cumsum_row_f32(src_spad_row, dst_spad_row, ne00);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
dma_queue_push_vtcm_to_ddr(dma_queue,
dma_make_ptr(dst_data + (ir * dst_row_size), (uint8_t *) dst_spad_row),
@@ -168,12 +172,10 @@ static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * dat
}
dma_queue_flush(dma_queue);
t2 = HAP_perf_get_qtimer_count();
FARF(HIGH, "cumsum-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
FARF(HIGH, "cumsum-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
}
// ---------------------------------------------------------------------------
@@ -183,14 +185,14 @@ static void cumsum_thread_f32_dma(unsigned int nth, unsigned int ith, void * dat
static void cumsum_thread_f32(unsigned int nth, unsigned int ith, void * data) {
htp_cumsum_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
const uint8_t * src_data = (const uint8_t *) src0->data;
uint8_t * dst_data = (uint8_t *) dst->data;
const uint32_t ir0 = cctx->rows_per_thread * ith;
const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->total_rows);
const uint32_t ir0 = cctx->row_start + cctx->rows_per_thread * ith;
const uint32_t ir1 = MIN(ir0 + cctx->rows_per_thread, cctx->row_start + cctx->total_rows);
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
for (uint32_t ir = ir0; ir < ir1; ir++) {
const float * restrict src_row = (const float *) (src_data + ir * cctx->src_row_size);
@@ -198,12 +200,11 @@ static void cumsum_thread_f32(unsigned int nth, unsigned int ith, void * data) {
hvx_cumsum_row_f32(src_row, dst_row, ne00);
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
FARF(HIGH, "cumsum-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
FARF(HIGH, "cumsum-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
}
int op_cumsum_f32(struct htp_ops_context * octx) {
@@ -214,8 +215,25 @@ int op_cumsum_f32(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads = MIN(octx->n_threads, total_rows);
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const size_t dst_data_row_size = dst->ne[0] * sizeof(float);
uint32_t row_start = 0;
uint32_t nrows = total_rows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_data_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_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 = octx->n_threads;
const size_t src_row_size = src0->nb[1];
const size_t dst_row_size = dst->nb[1];
@@ -240,14 +258,15 @@ int op_cumsum_f32(struct htp_ops_context * octx) {
.dst_row_size = dst_row_size,
.src_row_size_aligned = src_row_size_aligned,
.dst_row_size_aligned = dst_row_size_aligned,
.rows_per_thread = (total_rows + n_threads - 1) / n_threads,
.total_rows = total_rows,
.rows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
.total_rows = nrows,
.row_start = row_start,
};
if (octx->ctx->vtcm_size < spad_per_thread * n_threads) {
worker_pool_run_func(octx->ctx->worker_pool, cumsum_thread_f32, &cctx, n_threads);
work_queue_run(octx->ctx->work_queue, cumsum_thread_f32, &cctx, n_threads);
} else {
worker_pool_run_func(octx->ctx->worker_pool, cumsum_thread_f32_dma, &cctx, n_threads);
work_queue_run(octx->ctx->work_queue, cumsum_thread_f32_dma, &cctx, n_threads);
}
return HTP_STATUS_OK;
+66 -35
View File
@@ -5,8 +5,11 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#include "hvx-types.h"
#include "hex-utils.h"
#include "hvx-copy.h"
@@ -15,17 +18,17 @@
#define htp_diag_tensors_preamble \
const struct htp_tensor * restrict src0 = octx->src[0]; \
const struct htp_tensor * restrict dst = octx->dst; \
\
const uint32_t ne02 = src0->ne[2]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
\
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
\
const uint32_t ne02 = src0->ne[2]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
\
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3];
struct htp_diag_context {
@@ -36,6 +39,7 @@ struct htp_diag_context {
size_t dst_row_size_aligned;
uint32_t batches_per_thread;
uint32_t total_batches;
uint32_t batch_start;
};
#define htp_diag_preamble \
@@ -57,11 +61,8 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
htp_diag_preamble;
dma_queue * dma_queue = octx->ctx->dma[ith];
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
const uint32_t ib0 = dctx->batches_per_thread * ith;
const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->total_batches);
const uint32_t ib0 = dctx->batch_start + dctx->batches_per_thread * ith;
const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->batch_start + dctx->total_batches);
if (ib0 >= ib1) {
return;
@@ -79,6 +80,8 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
uint8_t * src_spad = octx->src0_spad.data + (ith * src_batch_size_aligned);
uint8_t * dst_spad = octx->dst_spad.data + (ith * dst_row_size_aligned);
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ib = ib0; ib < ib1; ib++) {
const uint32_t i3 = ib / ne02;
const uint32_t i2 = ib % ne02;
@@ -96,7 +99,9 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
for (uint32_t i1 = 0; i1 < ne1; i1++) {
// Compute row in VTCM
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (ib * ne1 + i1));
hvx_diag_row_f32(src_spad_f32, dst_spad_f32, i1, ne0);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (ib * ne1 + i1));
// Write completed row back to DDR
uint8_t * dst_row = dst_data + i3 * nb3 + i2 * nb2 + i1 * nb1;
@@ -107,12 +112,9 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
}
}
t2 = HAP_perf_get_qtimer_count();
FARF(HIGH, "diag-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
FARF(HIGH, "diag-f32-dma %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ib0, ib1,
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
}
// ---------------------------------------------------------------------------
@@ -122,14 +124,14 @@ static void diag_thread_f32_dma(unsigned int nth, unsigned int ith, void * data)
static void diag_thread_f32(unsigned int nth, unsigned int ith, void * data) {
htp_diag_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
const uint8_t * src_data = (const uint8_t *) src0->data;
uint8_t * dst_data = (uint8_t *) dst->data;
const uint32_t ib0 = dctx->batches_per_thread * ith;
const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->total_batches);
const uint32_t ib0 = dctx->batch_start + dctx->batches_per_thread * ith;
const uint32_t ib1 = MIN(ib0 + dctx->batches_per_thread, dctx->batch_start + dctx->total_batches);
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ib0);
for (uint32_t ib = ib0; ib < ib1; ib++) {
const uint32_t i3 = ib / ne02;
@@ -143,12 +145,11 @@ static void diag_thread_f32(unsigned int nth, unsigned int ith, void * data) {
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ib0);
FARF(HIGH, "diag-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u usec %u\n",
FARF(HIGH, "diag-f32 %d/%d: %ux%ux%ux%u (%u:%u) -> %ux%ux%ux%u\n",
ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ib0, ib1,
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
}
int op_diag_f32(struct htp_ops_context * octx) {
@@ -160,7 +161,36 @@ int op_diag_f32(struct htp_ops_context * octx) {
}
const uint32_t total_batches = src0->ne[2] * src0->ne[3];
const uint32_t n_threads = MIN(octx->n_threads, total_batches);
const size_t dst_batch_size = dst->ne[1] * dst->nb[1];
uint32_t batch_start = 0;
uint32_t nbatches = total_batches;
if (octx->ctx->mdev.count > 1) {
bool can_split = htp_tensor_mdev_data_aligned(dst) && (dst->ne[0] == 1 || dst->nb[0] == sizeof(float)) && !htp_tensor_is_permuted(dst);
uint32_t batches_per_chunk = 1;
if (can_split) {
if (dst->ne[2] > 1 && (dst->nb[2] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0 &&
(dst->ne[3] <= 1 || (dst->nb[3] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0)) {
batches_per_chunk = 1;
} else if (dst->nb[2] == dst_batch_size &&
(dst->ne[3] <= 1 || dst->nb[3] == dst->nb[2] * dst->ne[2])) {
batches_per_chunk = (dst_batch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(dst_batch_size, HEX_L2_LINE_SIZE)) : 1;
} else {
can_split = false;
}
}
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_batches, can_split ? batches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
batch_start = range.start;
nbatches = range.count;
}
if (nbatches == 0) {
return HTP_STATUS_OK;
}
const uint32_t n_threads = octx->n_threads;
const size_t src_batch_size = src0->ne[0] * sizeof(float);
const size_t dst_row_size = dst->ne[0] * sizeof(float);
@@ -185,14 +215,15 @@ int op_diag_f32(struct htp_ops_context * octx) {
.dst_row_size = dst_row_size,
.src_batch_size_aligned = src_batch_size_aligned,
.dst_row_size_aligned = dst_row_size_aligned,
.batches_per_thread = (total_batches + n_threads - 1) / n_threads,
.total_batches = total_batches,
.batches_per_thread = fastdiv(nbatches + n_threads - 1, &octx->n_threads_div),
.total_batches = nbatches,
.batch_start = batch_start,
};
if (octx->ctx->vtcm_size < spad_per_thread * n_threads) {
worker_pool_run_func(octx->ctx->worker_pool, diag_thread_f32, &dctx, n_threads);
work_queue_run(octx->ctx->work_queue, diag_thread_f32, &dctx, n_threads);
} else {
worker_pool_run_func(octx->ctx->worker_pool, diag_thread_f32_dma, &dctx, n_threads);
work_queue_run(octx->ctx->work_queue, diag_thread_f32_dma, &dctx, n_threads);
}
return HTP_STATUS_OK;
+46 -23
View File
@@ -3,10 +3,11 @@
#pragma clang diagnostic ignored "-Wunused-but-set-variable"
#include <HAP_farf.h>
#include <HAP_perf.h>
#include <string.h>
#include "hex-common.h"
#include "hex-profile.h"
#include "hvx-copy.h"
#include "hvx-utils.h"
@@ -14,28 +15,30 @@
#include "ggml-common.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
// ggml op_params layout for FILL:
// op_params[0] (as float) - the scalar fill value
#define fill_preamble \
#define fill_preamble \
const struct htp_tensor * dst = octx->dst; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3]; \
\
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3]; \
\
const uint32_t nr = ne1 * ne2 * ne3;
struct htp_fill_context {
struct htp_ops_context * octx;
uint32_t nrows_per_thread;
uint32_t total_rows; // ne1 * ne2 * ne3
uint32_t row_start;
bool opt_path;
HVX_Vector splat_vec;
uint32_t elem_size;
@@ -47,10 +50,15 @@ static void fill_thread(unsigned int nth, unsigned int ith, void * data) {
fill_preamble;
// Parallelise over the flat row index spanning ne1*ne2*ne3
const uint32_t ir0 = fctx->nrows_per_thread * ith;
const uint32_t ir1 = MIN(ir0 + fctx->nrows_per_thread, fctx->total_rows);
const uint32_t ir0 = fctx->row_start + fctx->nrows_per_thread * ith;
const uint32_t ir1 = MIN(ir0 + fctx->nrows_per_thread, fctx->row_start + fctx->total_rows);
uint64_t t1 = HAP_perf_get_qtimer_count();
if (ir0 >= ir1) {
return;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
if (fctx->opt_path) {
// Opt path: tensor is fully contiguous, treat as flat array
@@ -69,9 +77,8 @@ static void fill_thread(unsigned int nth, unsigned int ith, void * data) {
}
}
uint64_t t2 = HAP_perf_get_qtimer_count();
FARF(HIGH, "fill %u/%u: rows %u:%u usec %u\n",
ith, nth, ir0, ir1, (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir1);
FARF(HIGH, "fill %u/%u: rows %u:%u\n", ith, nth, ir0, ir1);
}
int op_fill(struct htp_ops_context * octx) {
@@ -85,8 +92,23 @@ int op_fill(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
uint32_t row_start = 0;
uint32_t nrows = nr;
if (octx->ctx->mdev.count > 1) {
const uint32_t row_size = nb1;
const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(nr, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0, 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;
}
// nr = ne1*ne2*ne3 (flat row count across all outer dims); parallelise over it.
const uint32_t n_threads = MIN(nr, octx->n_threads);
const uint32_t n_threads = octx->n_threads;
// Optimize if fully contiguous: skip stride arithmetic, treat as flat array
const bool opt_path = (nb2 == nb1 * ne1) && (nb3 == nb2 * ne2);
@@ -99,8 +121,9 @@ int op_fill(struct htp_ops_context * octx) {
struct htp_fill_context fctx = {
.octx = octx,
.nrows_per_thread = (nr + n_threads - 1) / n_threads,
.total_rows = nr,
.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
.total_rows = nrows,
.row_start = row_start,
.opt_path = opt_path,
};
@@ -117,7 +140,7 @@ int op_fill(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
worker_pool_run_func(octx->ctx->worker_pool, fill_thread, &fctx, n_threads);
work_queue_run(octx->ctx->work_queue, fill_thread, &fctx, n_threads);
return HTP_STATUS_OK;
}
+62 -22
View File
@@ -5,7 +5,6 @@
#include <assert.h>
#include <HAP_compute_res.h>
#include <HAP_farf.h>
#include <HAP_perf.h>
#include <math.h>
#include <stdbool.h>
#include <stdatomic.h>
@@ -75,6 +74,7 @@ struct htp_fa_context {
uint32_t qrows;
uint32_t qrows_per_thread;
uint32_t qrow_start;
bool is_q_fp32;
@@ -89,8 +89,6 @@ struct htp_fa_context {
const struct htp_tensor * k;
const struct htp_tensor * v;
uint64_t t_start;
};
struct hmx_fa_context {
@@ -206,10 +204,9 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
const uint32_t nb3 = dst->nb[3];
// total rows in q
const uint32_t nr = factx->qrows;
const uint32_t dr = factx->qrows_per_thread;
const uint32_t ir0 = dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, nr);
const uint32_t dr = factx->qrows_per_thread;
const uint32_t ir0 = factx->qrow_start + dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, factx->qrow_start + factx->qrows);
if (ir0 >= ir1) return;
@@ -1888,6 +1885,24 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
const uint32_t n_threads = factx.n_threads;
const uint32_t G = factx.G;
// Multi-device: split Q blocks across devices
const uint32_t n_q_blocks = (neq1 + Br - 1) / Br;
uint32_t q_start_min = 0;
uint32_t q_start_max = neq1;
if (octx->ctx->mdev.count > 1) {
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(n_q_blocks, htp_tensor_mdev_data_aligned(dst) ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
const uint32_t block_start = range.start;
const uint32_t block_end = range.start + range.count;
if (block_start >= block_end) {
return HTP_STATUS_OK;
}
q_start_min = block_start * Br;
q_start_max = MIN(block_end * Br, neq1);
}
// ======== VTCM allocation (GQA-aware) ========
// K/V row sizes drive the DMA descriptors (not the VTCM layout) and are used
// throughout the KV loop below.
@@ -1977,7 +1992,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
// ======== Main loop ========
for (uint32_t ib3 = 0; ib3 < neq3; ++ib3) {
const uint32_t im3 = mask ? fastmodulo(ib3, mask->ne[3], &factx.src3_div3) : 0;
for (uint32_t q_start = 0; q_start < neq1; q_start += Br) {
for (uint32_t q_start = q_start_min; q_start < q_start_max; q_start += Br) {
const uint32_t n_rows_q = hex_smin(Br, neq1 - q_start);
const size_t n_rows_g = n_rows_q * G;
const size_t g_br_actual = hex_align_up(n_rows_g, HMX_FP16_TILE_N_ROWS);
@@ -1991,8 +2006,9 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
// 1. Push Q and KV DMAs for the very first iteration.
// Subsequent iterations are enqueued early at the end of the previous iteration.
if (ib3 == 0 && q_start == 0 && kv_head == 0) {
const uint8_t * q_ptr = (const uint8_t *) q->data;
if (ib3 == 0 && q_start == q_start_min && kv_head == 0) {
const uint8_t * q_ptr = (const uint8_t *) q->data + q_start * q->nb[1] +
(kv_head * factx.G) * q->nb[2] + ib3 * q->nb[3];
const size_t q_row_bytes = q_transposed ? n_rows_q * q_row_bytes_trans_factor : q_row_bytes_untransposed;
const size_t n_rows = q_transposed ? factx.G : n_rows_q;
dma_queue_push(dma, dma_make_ptr(factx.vtcm_q_dma, q_ptr), q_row_bytes, hex_smax(q_src_stride, q_row_bytes), q_row_bytes, n_rows);
@@ -2311,8 +2327,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
if (next_kv_head >= n_kv_heads) {
next_kv_head = 0;
next_q_start = q_start + Br;
if (next_q_start >= neq1) {
next_q_start = 0;
if (next_q_start >= q_start_max) {
next_q_start = q_start_min;
next_ib3 = ib3 + 1;
}
}
@@ -2398,6 +2414,10 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
if (kparams->kernel_type == HTP_FA_KERNEL_HMX) {
return hmx_flash_attn_ext(octx);
}
@@ -2407,8 +2427,6 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
factx.k = k;
factx.v = v;
factx.t_start = HAP_perf_get_qtimer_count();
factx.src0_div21 = kparams->u.hvx.src0_div21;
factx.src0_div1 = kparams->u.hvx.src0_div1;
@@ -2451,8 +2469,30 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
}
// total rows in q
factx.qrows = kparams->qrows;
factx.qrows_per_thread = kparams->qrows_per_thread;
const uint32_t neq1 = q->ne[1];
const uint32_t neq2 = q->ne[2];
const uint32_t neq3 = q->ne[3];
const uint32_t total_qrows = neq1 * neq2 * neq3;
uint32_t qrow_start = 0;
uint32_t qrows = total_qrows;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) && ((dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_qrows, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
qrow_start = range.start;
qrows = range.count;
}
if (qrows == 0) {
return HTP_STATUS_OK;
}
const uint32_t n_threads = octx->n_threads;
factx.qrows = qrows;
factx.qrow_start = qrow_start;
factx.qrows_per_thread = fastdiv(qrows + n_threads - 1, &octx->n_threads_div);
size_t size_vkq_acc = hex_round_up(v->ne[0] * sizeof(float), 128); // VKQ32
@@ -2461,18 +2501,18 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
uint8_t * vtcm_cur = octx->ctx->vtcm_base;
factx.spad_q = vtcm_seq_alloc(&vtcm_cur, size_q_block * octx->n_threads);
factx.spad_k = vtcm_seq_alloc(&vtcm_cur, factx.size_k_block * 2 * octx->n_threads);
factx.spad_v = vtcm_seq_alloc(&vtcm_cur, factx.size_v_block * 2 * octx->n_threads);
factx.spad_m = vtcm_seq_alloc(&vtcm_cur, (mask ? factx.size_m_block * HVX_FA_DMA_CACHE_SIZE : 0) * octx->n_threads);
factx.spad_a = vtcm_seq_alloc(&vtcm_cur, size_vkq_acc * octx->n_threads);
factx.spad_q = vtcm_seq_alloc(&vtcm_cur, size_q_block * n_threads);
factx.spad_k = vtcm_seq_alloc(&vtcm_cur, factx.size_k_block * 2 * n_threads);
factx.spad_v = vtcm_seq_alloc(&vtcm_cur, factx.size_v_block * 2 * n_threads);
factx.spad_m = vtcm_seq_alloc(&vtcm_cur, (mask ? factx.size_m_block * HVX_FA_DMA_CACHE_SIZE : 0) * n_threads);
factx.spad_a = vtcm_seq_alloc(&vtcm_cur, size_vkq_acc * n_threads);
if ((size_t) (vtcm_cur - octx->ctx->vtcm_base) > octx->ctx->vtcm_size) {
return HTP_STATUS_VTCM_TOO_SMALL;
}
if (!(octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
work_queue_run(octx->ctx->work_queue, flash_attn_ext_f16_thread, &factx, octx->n_threads);
work_queue_run(octx->ctx->work_queue, flash_attn_ext_f16_thread, &factx, n_threads);
}
return HTP_STATUS_OK;
@@ -51,6 +51,7 @@ struct htp_fa_kernel_params {
uint32_t qrows;
uint32_t qrows_per_thread;
uint32_t qrow_start;
float m0;
float m1;
uint32_t n_head_log2;
+53 -21
View File
@@ -4,10 +4,13 @@
#include "hvx-utils.h"
#include "hex-fastdiv.h"
#include "hex-common.h"
#include "hex-profile.h"
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "htp-ctx.h"
#include "htp-tensor.h"
#ifndef MIN
#define MIN(a, b) ((a) < (b) ? (a) : (b))
@@ -22,6 +25,8 @@ struct htp_gdn_context {
size_t state_bytes;
uint8_t * vtcm_base;
size_t vtcm_per_thread;
uint32_t row_start;
uint32_t nrows;
};
static inline HVX_Vector gdn_mul_dot_f32(float * restrict dst, const float * restrict mul, const float * restrict dot, uint32_t n) {
@@ -586,8 +591,9 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
const uint32_t n_seqs = v->ne[3];
const uint32_t K = octx->op_params[0];
const uint32_t total_rows = H * n_seqs;
if (ith >= total_rows) {
const uint32_t row_end = gctx->row_start + gctx->nrows;
if (ith >= gctx->nrows) {
return;
}
@@ -621,11 +627,11 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
const uint64_t state_seq_stride = state->nb[3] / sizeof(float);
const uint64_t state_size_per_snap = (uint64_t) S_v * S_v * H * n_seqs;
uint32_t ir_prefetch = ith;
uint32_t ir_prefetch = gctx->row_start + ith;
int spad_idx = 0;
// Prefetch preamble (up to 2 steps)
for (int k = 0; k < 2 && ir_prefetch < total_rows; k++) {
for (int k = 0; k < 2 && ir_prefetch < row_end; k++) {
const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -646,8 +652,11 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (gctx->row_start + ith));
int curr_spad_idx = 0;
for (uint32_t ir = ith; ir < total_rows; ir += nth) {
for (uint32_t ir = gctx->row_start + ith; ir < row_end; ir += nth) {
dma_queue_pop(dma);
dma_queue_pop(dma);
@@ -812,7 +821,7 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
S_v * sizeof(float), S_v);
// Prefetch next block (if any)
if (ir_prefetch < total_rows) {
if (ir_prefetch < row_end) {
const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -828,6 +837,7 @@ static void gated_delta_net_f32_pp_thread(unsigned int nth, unsigned int ith, vo
curr_spad_idx ^= 1;
}
dma_queue_flush(dma);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) row_end);
}
@@ -847,8 +857,9 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
const uint32_t H = v->ne[1];
const uint32_t n_seqs = v->ne[3];
const uint32_t total_rows = H * n_seqs;
if (ith >= total_rows) {
const uint32_t row_end = gctx->row_start + gctx->nrows;
if (ith >= gctx->nrows) {
return;
}
@@ -881,11 +892,11 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
const uint64_t state_seq_stride = state->nb[3] / sizeof(float);
uint32_t ir_prefetch = ith;
uint32_t ir_prefetch = gctx->row_start + ith;
int spad_idx = 0;
// Prefetch preamble (up to 2 steps)
for (int k = 0; k < 2 && ir_prefetch < total_rows; k++) {
for (int k = 0; k < 2 && ir_prefetch < row_end; k++) {
const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -906,8 +917,11 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) (gctx->row_start + ith));
int curr_spad_idx = 0;
for (uint32_t ir = ith; ir < total_rows; ir += nth) {
for (uint32_t ir = gctx->row_start + ith; ir < row_end; ir += nth) {
dma_queue_pop(dma);
dma_queue_pop(dma);
@@ -1057,7 +1071,7 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
S_v * sizeof(float), S_v);
// Prefetch next block (if any)
if (ir_prefetch < total_rows) {
if (ir_prefetch < row_end) {
const uint32_t piv1 = fastmodulo(ir_prefetch, H, &fd_H);
const uint32_t piv3 = fastdiv(ir_prefetch, &fd_H);
const float * ps_in = state_in_base + (uint64_t) piv3 * state_seq_stride + (uint64_t) piv1 * S_v * S_v;
@@ -1073,6 +1087,7 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo
curr_spad_idx ^= 1;
}
dma_queue_flush(dma);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) row_end);
}
@@ -1085,10 +1100,6 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
const struct htp_tensor * state = octx->src[5];
const struct htp_tensor * dst = octx->dst;
if (!q || !k || !v || !g || !beta || !state || !dst) {
return HTP_STATUS_INVAL_PARAMS;
}
if (q->type != HTP_TYPE_F32 || k->type != HTP_TYPE_F32 || v->type != HTP_TYPE_F32 ||
g->type != HTP_TYPE_F32 || beta->type != HTP_TYPE_F32 || state->type != HTP_TYPE_F32 ||
dst->type != HTP_TYPE_F32) {
@@ -1124,16 +1135,37 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
const uint32_t total_rows = H * n_seqs;
uint32_t row_start = 0;
uint32_t nrows = total_rows;
if (octx->ctx->mdev.count > 1) {
const uint32_t head_bytes = S_v * sizeof(float);
const uint32_t rows_per_chunk = (head_bytes > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(head_bytes, HEX_L2_LINE_SIZE)) : 1;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0,
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 = octx->n_threads;
struct htp_gdn_context gctx;
gctx.octx = octx;
gctx.rows_per_thread = (H * n_seqs + octx->n_threads - 1) / octx->n_threads;
gctx.row_start = row_start;
gctx.nrows = nrows;
gctx.rows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
gctx.state_bytes = (size_t) S_v * S_v * sizeof(float);
size_t state_aligned = (size_t) S_v * S_v * sizeof(float);
state_aligned = (state_aligned + 127) & ~(size_t)127;
assert(octx->ctx->vtcm_base != NULL);
assert(octx->ctx->vtcm_size >= 2 * state_aligned * octx->n_threads);
assert(octx->ctx->vtcm_size >= 2 * state_aligned * n_threads);
gctx.vtcm_base = octx->ctx->vtcm_base;
gctx.vtcm_per_thread = 2 * state_aligned;
@@ -1148,9 +1180,9 @@ int op_gated_delta_net(struct htp_ops_context * octx) {
gctx.vtcm_per_thread * octx->n_threads, octx->n_threads);
if (n_tokens == 1) {
worker_pool_run_func(octx->ctx->worker_pool, gated_delta_net_f32_tg_thread, &gctx, octx->n_threads);
work_queue_run(octx->ctx->work_queue, gated_delta_net_f32_tg_thread, &gctx, n_threads);
} else {
worker_pool_run_func(octx->ctx->worker_pool, gated_delta_net_f32_pp_thread, &gctx, octx->n_threads);
work_queue_run(octx->ctx->work_queue, gated_delta_net_f32_pp_thread, &gctx, n_threads);
}
return HTP_STATUS_OK;
+47 -15
View File
@@ -10,6 +10,7 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
@@ -23,9 +24,12 @@ struct get_rows_context {
const struct htp_get_rows_kernel_params * kparams;
struct htp_get_rows_vtcm_layout vtcm_layout;
uint8_t * vtcm_base;
uint32_t task_start;
uint32_t tasks;
uint32_t tasks_per_thread;
};
#define get_rows_preamble \
#define get_rows_preamble \
const uint32_t ne00 = octx->src[0]->ne[0]; \
const uint32_t ne01 = octx->src[0]->ne[1]; \
const uint32_t ne02 = octx->src[0]->ne[2]; \
@@ -61,12 +65,12 @@ static void get_rows_thread_st_##IDX_TYPE(unsigned int nth, unsigned int ith, vo
struct htp_ops_context * octx = grctx->octx; \
const struct htp_get_rows_kernel_params * kparams = grctx->kparams; \
get_rows_preamble; \
const uint32_t dr = kparams->tasks_per_thread; \
const uint32_t ir0 = dr * ith; \
if (ir0 >= kparams->total_tasks) { \
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, kparams->total_tasks); \
const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks); \
const uint32_t row_size_bytes = htp_tensor_get_row_size(octx->src[0]->type, ne00); \
dma_queue * dma_queue = octx->ctx->dma[ith]; \
for (uint32_t i = ir0; i < ir1; ++i) { \
@@ -101,12 +105,12 @@ static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned
const struct htp_get_rows_kernel_params * kparams = grctx->kparams; \
get_rows_preamble; \
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
const uint32_t dr = kparams->tasks_per_thread; \
const uint32_t ir0 = dr * ith; \
if (ir0 >= kparams->total_tasks) { \
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, kparams->total_tasks); \
const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks); \
const uint32_t chunks_per_row = kparams->chunks_per_row; \
const uint32_t chunk_size = kparams->chunk_size; \
dma_queue * dma_queue = octx->ctx->dma[ith]; \
@@ -225,13 +229,41 @@ int op_get_rows(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
const struct htp_tensor * dst = octx->dst;
const uint32_t total_tasks = kparams->total_tasks;
const size_t dst_row_size = htp_tensor_get_row_size(dst->type, dst->ne[0]);
uint32_t task_start = 0;
uint32_t tasks = total_tasks;
if (octx->ctx->mdev.count > 1) {
uint32_t tasks_per_chunk = 1;
htp_tensor_mdev_rows_per_chunk(dst, dst_row_size / dst->ne[0], (uint32_t) dst_row_size, &tasks_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_tasks, tasks_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
task_start = range.start;
tasks = range.count;
}
if (tasks == 0) {
return HTP_STATUS_OK;
}
if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
const uint32_t n_threads = octx->n_threads;
struct get_rows_context grctx;
grctx.octx = octx;
grctx.kparams = kparams;
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);
const uint32_t ne00 = octx->src[0]->ne[0];
htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, octx->src[0]->type, ne00, kparams->n_threads);
htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, octx->src[0]->type, ne00, n_threads);
const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32);
@@ -247,14 +279,14 @@ int op_get_rows(struct htp_ops_context * octx) {
}
}
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 use-dma %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 * kparams->n_threads,
grctx.vtcm_layout.dst_bytes_per_thread * kparams->n_threads,
kparams->use_dma, kparams->n_threads);
grctx.vtcm_layout.src0_bytes_per_thread * n_threads,
grctx.vtcm_layout.dst_bytes_per_thread * n_threads,
kparams->use_dma, n_threads);
work_queue_run(octx->ctx->work_queue, q_func, &grctx, kparams->n_threads);
work_queue_run(octx->ctx->work_queue, q_func, &grctx, n_threads);
return HTP_STATUS_OK;
}
+9
View File
@@ -77,4 +77,13 @@ static inline bool hex_add_overflow(size_t a, size_t b, size_t *out) {
return false;
}
static inline uint32_t hex_gcd_u32(uint32_t a, uint32_t b) {
while (b != 0) {
uint32_t t = b;
b = a % b;
a = t;
}
return a;
}
#endif // HEX_COMMON_H
-1
View File
@@ -39,7 +39,6 @@ static inline void hex_l2fetch_block(const void * addr, size_t size) {
#define HEX_L2_LINE_SIZE 128
#define HEX_L2_BLOCK_SIZE (HEX_L2_LINE_SIZE * 4) // flush granularity (lines per loop iteration)
#define HEX_L2_FLUSH_IL_THRESHOLD 1024 // inline flush threshold
#define HEX_L2_FLUSH_WQ_THRESHOLD (4 * 1024)
#define HEX_L2_FLUSH_ALL_THRESHOLD (4 * 1024 * 1024)
+7 -7
View File
@@ -27,7 +27,7 @@ static inline void hmx_init_column_scales(void *out_scales, HVX_Vector v_scale)
// vscatter offsets for fused dequant+transpose: write K-values directly to [K][N] tile.
// word[i] = i*128 maps K-row-pair i to byte offset i*128.
// Column offset (n*4) is added at runtime. Entries 0..15 cover one tile (region 2047);
// entries 16..31 cover the next adjacent tile (region 4095) pick region size at the
// entries 16..31 cover the next adjacent tile (region 4095) - pick region size at the
// call site to scatter into one tile (masked) or two contiguous tiles (unmasked).
static const int32_t hmx_transpose_scatter_offsets[32] __attribute__((aligned(VLEN))) = {
0 * 128, 1 * 128, 2 * 128, 3 * 128, 4 * 128, 5 * 128, 6 * 128, 7 * 128, 8 * 128, 9 * 128, 10 * 128,
@@ -198,16 +198,16 @@ static inline void hmx_interleave_cols_to_tiles(__fp16 * restrict tiles_out,
}
// --- HMX inline asm macros for load-store packetization ---
#define HMX_LOAD_MPY_F16(act, wt, range) \
"{\n" \
#define HMX_LOAD_MPY_F16(act, wt, range) \
"{\n" \
" activation.hf = mxmem(" act ", " range ")\n" \
" weight.hf = mxmem(" wt ", " range ")\n" \
" weight.hf = mxmem(" wt ", " range ")\n" \
"}\n"
#define HMX_LOAD_MPY_DEEP_F16(act, wt, range) \
"{\n" \
#define HMX_LOAD_MPY_DEEP_F16(act, wt, range) \
"{\n" \
" activation.hf = mxmem(" act ", " range "):deep\n" \
" weight.hf = mxmem(" wt ", " range ")\n" \
" weight.hf = mxmem(" wt ", " range ")\n" \
"}\n"
#define HMX_STORE_AFTER_F16(out, scale_reg) \
+42 -9
View File
@@ -19,7 +19,7 @@
#endif
#define HTP_MAX_MMAPS 16
#define HTP_MAX_DIRTY_RANGES 16
#define HTP_MAX_DIRTY_RANGES 32
// Memory mapping
struct htp_mmap {
@@ -29,6 +29,11 @@ struct htp_mmap {
uint32_t reserved;
};
struct htp_dirty_range {
uint32_t start;
uint32_t end;
};
// Scratchpad state
struct htp_spad {
const struct htp_tensor * src; // original src of the data (for reuse)
@@ -38,6 +43,14 @@ struct htp_spad {
uint32_t size_per_thread; // size per thread
};
struct htp_mdev_group {
uint16_t idx;
uint16_t count;
struct fastdiv_values count_div;
uint8_t * fence_base;
uint32_t fence_seq;
};
struct htp_context;
// Context while processing an Op
@@ -65,8 +78,10 @@ struct htp_ops_context {
struct htp_spad src3_spad;
struct htp_spad dst_spad;
uint32_t n_threads;
uint32_t flags;
uint32_t flags;
uint32_t n_threads;
struct fastdiv_values n_threads_div;
int status;
};
// Main context for htp DSP backend
@@ -76,6 +91,7 @@ struct htp_context {
struct htp_mmap mmap[HTP_MAX_MMAPS];
dma_queue_t dma[HTP_MAX_NTHREADS];
dma_queue_t dma_cached[HTP_MAX_NTHREADS];
struct htp_thread_trace trace[HTP_MAX_NTHREADS + 1];
work_queue_t work_queue;
hmx_queue_t hmx_queue;
@@ -88,7 +104,6 @@ struct htp_context {
bool hmx_enabled;
bool etm;
uint32_t profiler;
struct htp_thread_trace trace[HTP_MAX_NTHREADS + 1];
uint8_t * vtcm_base;
size_t vtcm_size;
@@ -97,16 +112,13 @@ struct htp_context {
atomic_bool vtcm_needs_release;
uint64_t max_vmem;
struct htp_dirty_range {
uint32_t start;
uint32_t end;
uint32_t bi;
} dirty_ranges[HTP_MAX_DIRTY_RANGES];
struct htp_dirty_range dirty_ranges[HTP_MAX_DIRTY_RANGES];
// Persistent DDR scratchpad for MUL_MAT_ID mappings
void * ddr_spad_base;
size_t ddr_spad_size;
struct htp_mdev_group mdev;
struct htp_ops_context octx;
qurt_thread_t main_thread;
@@ -115,6 +127,27 @@ struct htp_context {
size_t footprint;
};
static inline bool htp_ops_context_set_n_threads(struct htp_ops_context * octx, uint32_t n_threads) {
if (n_threads == 0 || n_threads > octx->ctx->n_threads) {
return false;
}
if (n_threads != octx->n_threads) {
octx->n_threads = n_threads;
octx->n_threads_div = n_threads == octx->ctx->n_threads
? octx->ctx->n_threads_div
: init_fastdiv_values(n_threads);
}
return true;
}
static inline void htp_ops_context_set_status(struct htp_ops_context * octx, int status) {
if (status > HTP_STATUS_OK && octx->status == HTP_STATUS_OK) {
octx->status = status;
}
}
int op_matmul(struct htp_ops_context * octx);
int op_matmul_id(struct htp_ops_context * octx);
int op_matmul_nx(struct htp_ops_context * octx);
+89
View File
@@ -0,0 +1,89 @@
#ifndef HTP_FENCE_H
#define HTP_FENCE_H
#include <stdatomic.h>
#include <stdint.h>
#include <HAP_farf.h>
#include "hex-utils.h"
#include "htp-ops.h"
#include "htp-ctx.h"
static inline atomic_uint * htp_mdev_fence_slot(const void * fence_base, uint32_t idx) {
return (atomic_uint *) ((const uint8_t *) fence_base + (size_t) idx * HTP_FENCE_SLOT_SIZE);
}
static inline void htp_fence_write(void * fence_ptr, uint32_t seq, uint32_t status) {
atomic_uint * fence = (atomic_uint *) fence_ptr;
atomic_store(&fence[1], status);
atomic_store(&fence[0], seq);
asm volatile ("syncht" : : : "memory");
Q6_dccleaninva_A((void *) fence);
}
static inline void htp_fence_read(const void * fence_ptr, uint32_t * seq, uint32_t * status) {
const atomic_uint * fence = (const atomic_uint *) fence_ptr;
Q6_dccleaninva_A((void *) fence);
asm volatile ("syncht" : : : "memory");
*seq = atomic_load(&fence[0]);
*status = atomic_load(&fence[1]);
}
static inline void htp_mdev_group_barrier(struct htp_ops_context * octx) {
struct htp_context * ctx = octx->ctx;
if (ctx->mdev.count <= 1) {
return;
}
const uint32_t seq = ++ctx->mdev.fence_seq;
struct htp_thread_trace * tr = &ctx->trace[0];
htp_trace_event_start(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
const uint32_t mdev_idx = ctx->mdev.idx;
const uint32_t mdev_count = ctx->mdev.count;
uint8_t * fence_base = ctx->mdev.fence_base;
atomic_uint * my_fence = htp_mdev_fence_slot(fence_base, mdev_idx);
htp_fence_write(my_fence, seq, octx->status);
for (uint32_t d = 0; d < mdev_count; d++) {
if (d == mdev_idx) continue;
atomic_uint * peer_fence = htp_mdev_fence_slot(fence_base, d);
uint64_t spins = 0;
while (1) {
uint32_t peer_seq;
uint32_t peer_status;
htp_fence_read(peer_fence, &peer_seq, &peer_status);
if ((int32_t)(peer_seq - seq) >= 0) {
if (peer_status > HTP_STATUS_OK) {
FARF(ERROR, "ggml-hex: mdev %u peer %u failed with status %u : seq 0x%08x\n",
mdev_idx, d, peer_status, seq);
htp_ops_context_set_status(octx, peer_status);
}
break;
}
if (++spins == 10000) {
FARF(ALWAYS, "ggml-hex: mdev %u waiting for mdev %u : seq 0x%08x (b %u op %u) my-fence %p peer-fence %p peer-seq 0x%08x (diff %d)\n",
mdev_idx, d, seq, seq >> 12, seq & 0xfff, my_fence, peer_fence, peer_seq, (int32_t)(peer_seq - seq));
}
if (spins > HTP_FENCE_TIMEOUT) {
FARF(ERROR, "ggml-hex: mdev %u timeout waiting for mdev %u : seq 0x%08x (b %u op %u) peer-fence %p peer-seq 0x%08x\n",
mdev_idx, d, seq, seq >> 12, seq & 0xfff, peer_fence, peer_seq);
htp_ops_context_set_status(octx, HTP_STATUS_INTERNAL_ERR);
break;
}
hex_pause();
}
}
asm volatile ("syncht" : : : "memory");
if (octx->status > HTP_STATUS_OK) {
htp_fence_write(my_fence, seq, octx->status);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
}
#endif // HTP_FENCE_H
+11 -12
View File
@@ -77,6 +77,7 @@ enum htp_op_code {
HTP_OP_GET_ROWS,
HTP_OP_SCALE,
HTP_OP_CPY,
HTP_OP_CPY_FENCE,
HTP_OP_ARGSORT,
HTP_OP_SQR,
HTP_OP_SQRT,
@@ -100,6 +101,7 @@ enum htp_op_code {
HTP_OP_ALLREDUCE,
HTP_OP_ALLREDUCE_ADD,
HTP_OP_GLU_SWIGLU_CLAMP,
HTP_OP_MDEV_GROUP,
HTP_OP_INVALID
};
@@ -114,6 +116,7 @@ enum htp_op_code {
#define HTP_OP_MAX_TENSORS 8192 // must stay under 64K (uint16)
#define HTP_FENCE_TIMEOUT (1000000000ULL)
#define HTP_FENCE_SLOT_SIZE 128
#define HTP_OP_MAX_VMEM_DEFAULT (3355443200u)
@@ -214,30 +217,26 @@ struct htp_prof_desc {
};
struct htp_opbatch_req {
uint32_t id; // Batch id
uint64_t seq; // Sequence number
uint32_t n_bufs; // Number of buffers
uint32_t n_tensors; // Number of tensors
uint32_t n_ops; // Number of ops
uint32_t n_traces; // Number of trace descriptors per thread
uint32_t pad; // unused
uint64_t seq; // Sequence number
// struct htp_buf_desc bufs[]; -- dspqueue buf 0
// struct htp_tensor tensors[]; -- dspqueue buf 0
// struct htp_op_desc ops[]; -- dspqueue buf 0
};
struct htp_opbatch_rsp {
uint32_t id; // Batch id
uint32_t status; // HTP_STATUS_...
uint32_t n_bufs; // Number of buffers
uint32_t n_tensors; // Number of tensors
uint32_t n_ops; // Number of op profile descriptors
uint32_t n_traces[HTP_MAX_NTHREADS + 1];
uint32_t usecs; // Number of usec
uint32_t pad; // align to 8 bytes
uint64_t seq; // Sequence number
uint64_t cycles_start; // Start cycle counter
uint64_t cycles_stop; // Stop cycle counter
uint64_t seq; // Sequence number
uint32_t status; // HTP_STATUS_...
uint32_t n_bufs; // Number of buffers
uint32_t n_tensors; // Number of tensors
uint32_t n_ops; // Number of op profile descriptors
uint32_t usecs; // Number of usec
uint32_t n_traces[HTP_MAX_NTHREADS + 1];
// struct htp_prof_desc profs[]; -- dspqueue buf 0
};
+94 -43
View File
@@ -20,7 +20,7 @@ struct l2flush_range {
struct l2flush_multi_task {
struct htp_thread_trace * trace;
struct l2flush_range ranges[HTP_OP_MAX_INPUTS];
struct l2flush_range ranges[HTP_MAX_DIRTY_RANGES];
uint32_t n_ranges;
uint32_t total_blocks;
uint32_t blocks_per_thread;
@@ -73,6 +73,27 @@ static void l2flush_multi_worker(unsigned int n, unsigned int i, void * data) {
htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, gb_first);
}
static void merge_dirty_ranges(struct htp_context * ctx) {
for (uint32_t i = 0; i < HTP_MAX_DIRTY_RANGES; i++) {
struct htp_dirty_range * r = &ctx->dirty_ranges[i];
if (!r->start) continue;
for (uint32_t j = 0; j < HTP_MAX_DIRTY_RANGES;) {
struct htp_dirty_range * s = &ctx->dirty_ranges[j];
if (i == j || !s->start || r->end < s->start || s->end < r->start) {
j++;
continue;
}
r->start = MIN(r->start, s->start);
r->end = MAX(r->end, s->end);
s->start = 0;
s->end = 0;
j = 0;
}
}
}
void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) {
const struct htp_tensor * pending[HTP_OP_MAX_OUTPUTS];
uint32_t n_pending = 0;
@@ -83,11 +104,6 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
continue;
}
if (t->size <= HEX_L2_FLUSH_IL_THRESHOLD) {
hex_l2flush((void *) (uintptr_t) t->data, t->size);
continue;
}
uint32_t t_start = t->data;
uint32_t t_end = t_start + t->size;
@@ -110,6 +126,8 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
}
}
merge_dirty_ranges(ctx);
if (n_pending == 0) {
return;
}
@@ -132,8 +150,8 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
struct htp_dirty_range * r = &ctx->dirty_ranges[idx];
r->start = pending[i]->data;
r->end = pending[i]->data + pending[i]->size;
r->bi = pending[i]->bi;
}
merge_dirty_ranges(ctx);
return;
}
@@ -151,12 +169,12 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
struct htp_dirty_range * r = &ctx->dirty_ranges[i];
r->start = pending[i]->data;
r->end = pending[i]->data + pending[i]->size;
r->bi = pending[i]->bi;
}
merge_dirty_ranges(ctx);
return;
}
if (total_evict_size > HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1 && n_evict <= HTP_OP_MAX_INPUTS) {
if (total_evict_size > HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1 && n_evict <= HTP_MAX_DIRTY_RANGES) {
struct l2flush_multi_task task;
task.trace = ctx->trace;
task.n_ranges = n_evict;
@@ -195,7 +213,6 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
struct htp_dirty_range * r = &ctx->dirty_ranges[idx];
r->start = pending[i]->data;
r->end = pending[i]->data + pending[i]->size;
r->bi = pending[i]->bi;
}
for (uint32_t i = 0; i < n_empty; i++) {
@@ -203,8 +220,9 @@ void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * co
struct htp_dirty_range * r = &ctx->dirty_ranges[idx];
r->start = pending[n_evict + i]->data;
r->end = pending[n_evict + i]->data + pending[n_evict + i]->size;
r->bi = pending[n_evict + i]->bi;
}
merge_dirty_ranges(ctx);
}
static void make_tensor_clean(struct htp_context * ctx, const struct htp_tensor * t) {
@@ -242,15 +260,77 @@ static inline bool is_tensor_dirty(struct htp_context * ctx, const struct htp_te
return false;
}
static void flush_dirty_ranges(struct htp_context * ctx, const struct htp_dirty_range * ranges, uint32_t n_ranges, uint64_t total_dirty) {
if (total_dirty >= HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1) {
struct l2flush_multi_task task;
task.trace = ctx->trace;
task.n_ranges = n_ranges;
uint32_t block_acc = 0;
for (uint32_t i = 0; i < n_ranges; i++) {
const struct htp_dirty_range * r = &ranges[i];
struct l2flush_range * rg = &task.ranges[i];
rg->start = hex_align_down((size_t) r->start, HEX_L2_LINE_SIZE);
rg->end = hex_align_up((size_t) r->end, HEX_L2_LINE_SIZE);
rg->block_first = block_acc;
rg->n_blocks = (rg->end - rg->start + HEX_L2_BLOCK_SIZE - 1) / HEX_L2_BLOCK_SIZE;
block_acc += rg->n_blocks;
}
task.total_blocks = block_acc;
task.blocks_per_thread = fastdiv(block_acc + ctx->n_threads - 1, &ctx->n_threads_div);
work_queue_run(ctx->work_queue, l2flush_multi_worker, &task, ctx->n_threads);
} else {
struct htp_thread_trace * tr = &ctx->trace[0];
htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, 0);
for (uint32_t i = 0; i < n_ranges; i++) {
const struct htp_dirty_range * r = &ranges[i];
hex_l2flush((void *) (uintptr_t) r->start, r->end - r->start);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, 0);
}
}
void htp_flush_dirty_ranges(struct htp_context * ctx) {
struct htp_dirty_range ranges[HTP_MAX_DIRTY_RANGES];
uint32_t n_ranges = 0;
uint64_t total_dirty = 0;
for (uint32_t i = 0; i < HTP_MAX_DIRTY_RANGES; i++) {
const struct htp_dirty_range * r = &ctx->dirty_ranges[i];
if (!r->start) {
continue;
}
ranges[n_ranges++] = *r;
total_dirty += r->end - r->start;
}
if (total_dirty == 0) {
return;
}
if (total_dirty > HEX_L2_FLUSH_ALL_THRESHOLD) {
flush_all_dcache(ctx);
return;
}
flush_dirty_ranges(ctx, ranges, n_ranges, total_dirty);
memset(ctx->dirty_ranges, 0, sizeof(ctx->dirty_ranges));
}
void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n) {
const struct htp_tensor * dirty_tensors[HTP_OP_MAX_INPUTS];
struct htp_dirty_range ranges[HTP_OP_MAX_INPUTS];
uint32_t n_dirty = 0;
uint64_t total_dirty = 0;
for (uint32_t i = 0; i < n; i++) {
const struct htp_tensor * t = tensors[i];
if (t && !(t->flags & (HTP_TENSOR_WEIGHT | HTP_TENSOR_FENCE)) && is_tensor_dirty(ctx, t)) {
if (t && is_tensor_dirty(ctx, t)) {
dirty_tensors[n_dirty++] = t;
ranges[n_dirty - 1].start = t->data;
ranges[n_dirty - 1].end = t->data + t->size;
total_dirty += t->size;
}
}
@@ -264,37 +344,8 @@ void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * co
return;
}
if (total_dirty >= HEX_L2_FLUSH_WQ_THRESHOLD && ctx->n_threads > 1) {
struct l2flush_multi_task task;
task.trace = ctx->trace;
task.n_ranges = 0;
uint32_t block_acc = 0;
for (uint32_t i = 0; i < n_dirty; i++) {
const struct htp_tensor * t = dirty_tensors[i];
make_tensor_clean(ctx, t);
struct l2flush_range * rg = &task.ranges[task.n_ranges++];
rg->start = hex_align_down((size_t) t->data, HEX_L2_LINE_SIZE);
rg->end = hex_align_up((size_t) t->data + t->size, HEX_L2_LINE_SIZE);
rg->block_first = block_acc;
rg->n_blocks = (rg->end - rg->start + HEX_L2_BLOCK_SIZE - 1) / HEX_L2_BLOCK_SIZE;
block_acc += rg->n_blocks;
}
task.total_blocks = block_acc;
task.blocks_per_thread = fastdiv(block_acc + ctx->n_threads - 1, &ctx->n_threads_div);
work_queue_run(ctx->work_queue, l2flush_multi_worker, &task, ctx->n_threads);
return;
}
struct htp_thread_trace * tr = &ctx->trace[0];
flush_dirty_ranges(ctx, ranges, n_dirty, total_dirty);
for (uint32_t i = 0; i < n_dirty; i++) {
const struct htp_tensor * t = dirty_tensors[i];
htp_trace_event_start(tr, HTP_TRACE_EVT_L2FLUSH, t->ti);
hex_l2flush((void *) (uintptr_t) t->data, t->size);
htp_trace_event_stop(tr, HTP_TRACE_EVT_L2FLUSH, t->ti);
make_tensor_clean(ctx, t);
make_tensor_clean(ctx, dirty_tensors[i]);
}
}
+109
View File
@@ -2,8 +2,20 @@
#define HTP_TENSOR_H
#include <stdint.h>
#include <stdbool.h>
#include "htp-ops.h"
#include "hex-bitmap.h"
#include "hex-common.h"
#include "hex-fastdiv.h"
enum {
HTP_TENSOR_MDEV_LINE_SIZE = 128,
};
struct htp_tensor_mdev_range {
uint32_t start;
uint32_t count;
};
static inline void * htp_tensor_data(const struct htp_tensor * t) {
return (void *) (uintptr_t) t->data;
@@ -13,6 +25,102 @@ static inline uint32_t * htp_tensor_flags(const struct htp_tensor * t) {
return (uint32_t *) &t->flags;
}
static inline bool htp_tensor_is_contiguous(const struct htp_tensor * t, uint32_t type_size) {
uint32_t next_nb = type_size;
if (t->ne[0] != 1 && t->nb[0] != next_nb) {
return false;
}
next_nb *= t->ne[0];
for (int i = 1; i < HTP_OP_MAX_DIMS; i++) {
if (t->ne[i] != 1 && t->nb[i] != next_nb) {
return false;
}
next_nb *= t->ne[i];
}
return true;
}
static inline bool htp_tensor_is_permuted(const struct htp_tensor * t) {
return t->nb[0] > t->nb[1] || t->nb[1] > t->nb[2] || t->nb[2] > t->nb[3];
}
static inline bool htp_tensor_mdev_data_aligned(const struct htp_tensor * t) {
return ((uintptr_t) t->data & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0;
}
static inline bool htp_tensor_can_row_partition(const struct htp_tensor * t, uint32_t elem_size) {
if (!htp_tensor_mdev_data_aligned(t)) {
return false;
}
if (t->ne[0] != 1 && t->nb[0] != elem_size) {
return false;
}
if (htp_tensor_is_permuted(t)) {
return false;
}
if (t->ne[1] > 1 && (t->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) != 0) return false;
if (t->ne[2] > 1 && (t->nb[2] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) != 0) return false;
if (t->ne[3] > 1 && (t->nb[3] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) != 0) return false;
return true;
}
static inline bool htp_tensor_mdev_rows_per_chunk(const struct htp_tensor * t, uint32_t elem_size, uint32_t row_size, uint32_t * rows_per_chunk) {
*rows_per_chunk = 0;
if (!htp_tensor_mdev_data_aligned(t)) {
return false;
}
if (t->ne[0] != 1 && t->nb[0] != elem_size) {
return false;
}
if (htp_tensor_is_permuted(t)) {
return false;
}
if (t->ne[1] > 1 && (t->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0 &&
(t->ne[2] <= 1 || (t->nb[2] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0) &&
(t->ne[3] <= 1 || (t->nb[3] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0)) {
*rows_per_chunk = 1;
return true;
}
if (t->nb[1] == row_size &&
(t->ne[2] <= 1 || t->nb[2] == t->nb[1] * t->ne[1]) &&
(t->ne[3] <= 1 || t->nb[3] == t->nb[2] * t->ne[2])) {
*rows_per_chunk = (row_size > 0) ? (HTP_TENSOR_MDEV_LINE_SIZE / hex_gcd_u32(row_size, HTP_TENSOR_MDEV_LINE_SIZE)) : 1;
return true;
}
return false;
}
static inline struct htp_tensor_mdev_range htp_tensor_mdev_partition(uint32_t total_units, uint32_t units_per_chunk, uint32_t mdev_idx, uint32_t mdev_count, const struct fastdiv_values * mdev_count_div) {
struct htp_tensor_mdev_range range = { 0, total_units };
if (mdev_count <= 1) {
return range;
}
if (units_per_chunk == 0) {
range.start = (mdev_idx == 0) ? 0 : total_units;
range.count = (mdev_idx == 0) ? total_units : 0;
return range;
}
const uint32_t total_chunks = total_units / units_per_chunk;
if (total_chunks < mdev_count) {
range.start = (mdev_idx == 0) ? 0 : total_units;
range.count = (mdev_idx == 0) ? total_units : 0;
return range;
}
const uint32_t chunks_per_mdev = fastdiv(total_chunks + mdev_count - 1, mdev_count_div);
range.start = MIN(mdev_idx * chunks_per_mdev * units_per_chunk, total_units);
if (mdev_idx == mdev_count - 1) {
range.count = total_units - range.start;
} else {
range.count = MIN(chunks_per_mdev * units_per_chunk, total_units - range.start);
}
return range;
}
static inline uint32_t htp_tensor_get_row_size(int type, uint32_t ne00) {
switch (type) {
case HTP_TYPE_F32: return ne00 * 4;
@@ -23,6 +131,7 @@ static inline uint32_t htp_tensor_get_row_size(int type, uint32_t ne00) {
}
struct htp_context;
void htp_flush_dirty_ranges(struct htp_context * ctx);
void htp_tensor_flush_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n);
void htp_tensor_dirty_all(struct htp_context * ctx, const struct htp_tensor * const * tensors, uint32_t n);
+137 -137
View File
@@ -16,25 +16,25 @@
#define UNUSED(x) (void)(x)
#define hvx_arith_loop_body(dst_type, src0_type, src1_type, elem_size, vec_store, vec_op) \
do { \
dst_type * vdst = (dst_type *) dst; \
src0_type * vsrc0 = (src0_type *) src0; \
src1_type * vsrc1 = (src1_type *) src1; \
\
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] = vec_op(vsrc0[i], vsrc1[i]); \
} \
if (nloe) { \
HVX_Vector v = vec_op(vsrc0[i], vsrc1[i]); \
vec_store((void *) &vdst[i], nloe * (elem_size), v); \
} \
do { \
dst_type * vdst = (dst_type *) dst; \
src0_type * vsrc0 = (src0_type *) src0; \
src1_type * vsrc1 = (src1_type *) src1; \
\
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] = vec_op(vsrc0[i], vsrc1[i]); \
} \
if (nloe) { \
HVX_Vector v = vec_op(vsrc0[i], vsrc1[i]); \
vec_store((void *) &vdst[i], nloe * (elem_size), v); \
} \
} while(0)
#if __HVX_ARCH__ < 79
@@ -56,43 +56,43 @@
#define HVX_OP_MUL_F16(a, b) hvx_vec_mul_f16_f16(a, b)
// Generic macro to define alignment permutations for an op
#define DEFINE_HVX_BINARY_OP_VARIANTS(OP_NAME, OP_MACRO, ELEM_TYPE) \
static inline void OP_NAME##_aaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_aau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_aua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_auu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_uaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
static inline void OP_NAME##_uau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) src0 % 128 == 0); \
hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
static inline void OP_NAME##_uua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
static inline void OP_NAME##_uuu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
#define DEFINE_HVX_BINARY_OP_VARIANTS(OP_NAME, OP_MACRO, ELEM_TYPE) \
static inline void OP_NAME##_aaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_aau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_aua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_auu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
hvx_arith_loop_body(HVX_Vector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_uaa(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
static inline void OP_NAME##_uau(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) src0 % 128 == 0); \
hvx_arith_loop_body(HVX_UVector, HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
static inline void OP_NAME##_uua(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
assert((uintptr_t) src1 % 128 == 0); \
hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
static inline void OP_NAME##_uuu(uint8_t * dst, const uint8_t * src0, const uint8_t * src1, uint32_t n) { \
hvx_arith_loop_body(HVX_UVector, HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
} \
DEFINE_HVX_BINARY_OP_VARIANTS(hvx_add_f32, HVX_OP_ADD_F32, float)
DEFINE_HVX_BINARY_OP_VARIANTS(hvx_sub_f32, HVX_OP_SUB_F32, float)
@@ -103,25 +103,25 @@ DEFINE_HVX_BINARY_OP_VARIANTS(hvx_sub_f16, HVX_OP_SUB_F16, _Float16)
DEFINE_HVX_BINARY_OP_VARIANTS(hvx_mul_f16, HVX_OP_MUL_F16, _Float16)
// Dispatcher logic
#define HVX_BINARY_DISPATCHER(OP_NAME) \
#define HVX_BINARY_DISPATCHER(OP_NAME) \
static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, const uint32_t num_elems) { \
if (hex_is_aligned((void *) dst, 128)) { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems); \
else OP_NAME##_aau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems); \
else OP_NAME##_auu(dst, src0, src1, num_elems); \
} \
} else { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems); \
else OP_NAME##_uau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems); \
else OP_NAME##_uuu(dst, src0, src1, num_elems); \
} \
} \
if (hex_is_aligned((void *) dst, 128)) { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems); \
else OP_NAME##_aau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems); \
else OP_NAME##_auu(dst, src0, src1, num_elems); \
} \
} else { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems); \
else OP_NAME##_uau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems); \
else OP_NAME##_uuu(dst, src0, src1, num_elems); \
} \
} \
}
HVX_BINARY_DISPATCHER(hvx_add_f32)
@@ -166,44 +166,44 @@ static inline void hvx_mul_mul_f32_aa(uint8_t * restrict dst, const uint8_t * re
// Scalar Operations
#define hvx_scalar_loop_body(dst_type, src_type, elem_size, vec_store, scalar_op_macro) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
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++) { \
HVX_Vector v = vsrc[i]; \
vdst[i] = scalar_op_macro(v); \
} \
if (nloe) { \
HVX_Vector v = vsrc[i]; \
v = scalar_op_macro(v); \
vec_store((void *) &vdst[i], nloe * (elem_size), v); \
} \
#define hvx_scalar_loop_body(dst_type, src_type, elem_size, vec_store, scalar_op_macro) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
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++) { \
HVX_Vector v = vsrc[i]; \
vdst[i] = scalar_op_macro(v); \
} \
if (nloe) { \
HVX_Vector v = vsrc[i]; \
v = scalar_op_macro(v); \
vec_store((void *) &vdst[i], nloe * (elem_size), v); \
} \
} while(0)
#define HVX_OP_ADD_SCALAR_F32(v) \
({ \
#define HVX_OP_ADD_SCALAR_F32(v) \
({ \
const HVX_VectorPred pred_inf = Q6_Q_vcmp_eq_VwVw(inf, v); \
HVX_Vector out = HVX_OP_ADD_F32(v, val_vec); \
Q6_V_vmux_QVV(pred_inf, inf, out); \
HVX_Vector out = HVX_OP_ADD_F32(v, val_vec); \
Q6_V_vmux_QVV(pred_inf, inf, out); \
})
#define HVX_OP_MUL_SCALAR_F32(v) HVX_OP_MUL_F32(v, val_vec)
#define HVX_OP_SUB_SCALAR_F32(v) HVX_OP_SUB_F32(v, val_vec)
#define HVX_OP_ADD_SCALAR_F16(v) \
({ \
#define HVX_OP_ADD_SCALAR_F16(v) \
({ \
const HVX_VectorPred pred_inf = Q6_Q_vcmp_eq_VhVh(inf, v); \
HVX_Vector out = HVX_OP_ADD_F16(v, val_vec); \
Q6_V_vmux_QVV(pred_inf, inf, out); \
HVX_Vector out = HVX_OP_ADD_F16(v, val_vec); \
Q6_V_vmux_QVV(pred_inf, inf, out); \
})
#define HVX_OP_MUL_SCALAR_F16(v) HVX_OP_MUL_F16(v, val_vec)
@@ -212,31 +212,31 @@ static inline void hvx_mul_mul_f32_aa(uint8_t * restrict dst, const uint8_t * re
// Scalar Variants
// Generic macro to define alignment permutations for an op
#define DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(OP_NAME, OP_MACRO, SPLAT_MACRO, ELEM_TYPE) \
#define DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(OP_NAME, OP_MACRO, SPLAT_MACRO, ELEM_TYPE) \
static inline void OP_NAME##_aa(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src % 128 == 0); \
hvx_scalar_loop_body(HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src % 128 == 0); \
hvx_scalar_loop_body(HVX_Vector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_au(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
assert((uintptr_t) dst % 128 == 0); \
hvx_scalar_loop_body(HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
assert((uintptr_t) dst % 128 == 0); \
hvx_scalar_loop_body(HVX_Vector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_a, OP_MACRO); \
} \
static inline void OP_NAME##_ua(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
assert((uintptr_t) src % 128 == 0); \
hvx_scalar_loop_body(HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
assert((uintptr_t) src % 128 == 0); \
hvx_scalar_loop_body(HVX_UVector, HVX_Vector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
static inline void OP_NAME##_uu(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, uint32_t n) { \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
hvx_scalar_loop_body(HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
const HVX_Vector val_vec = SPLAT_MACRO(val); \
const HVX_Vector inf = SPLAT_MACRO((ELEM_TYPE)INFINITY); UNUSED(inf); \
hvx_scalar_loop_body(HVX_UVector, HVX_UVector, sizeof(ELEM_TYPE), hvx_vec_store_u, OP_MACRO); \
} \
DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_add_scalar_f32, HVX_OP_ADD_SCALAR_F32, hvx_vec_splat_f32, float)
DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_sub_scalar_f32, HVX_OP_SUB_SCALAR_F32, hvx_vec_splat_f32, float)
@@ -247,17 +247,17 @@ DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_sub_scalar_f16, HVX_OP_SUB_SCALAR_F16,
DEFINE_HVX_BINARY_SCALAR_OP_VARIANTS(hvx_mul_scalar_f16, HVX_OP_MUL_SCALAR_F16, hvx_vec_splat_f16, _Float16)
// Dispatcher logic
#define HVX_BINARY_SCALAR_DISPATCHER(OP_NAME, ELEM_TYPE) \
#define HVX_BINARY_SCALAR_DISPATCHER(OP_NAME, ELEM_TYPE) \
static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src, const ELEM_TYPE val, const uint32_t num_elems) { \
if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) { \
OP_NAME##_aa(dst, src, val, num_elems); \
} else if (hex_is_aligned((void *) dst, 128)) { \
OP_NAME##_au(dst, src, val, num_elems); \
} else if (hex_is_aligned((void *) src, 128)) { \
OP_NAME##_ua(dst, src, val, num_elems); \
} else { \
OP_NAME##_uu(dst, src, val, num_elems); \
} \
if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) { \
OP_NAME##_aa(dst, src, val, num_elems); \
} else if (hex_is_aligned((void *) dst, 128)) { \
OP_NAME##_au(dst, src, val, num_elems); \
} else if (hex_is_aligned((void *) src, 128)) { \
OP_NAME##_ua(dst, src, val, num_elems); \
} else { \
OP_NAME##_uu(dst, src, val, num_elems); \
} \
}
HVX_BINARY_SCALAR_DISPATCHER(hvx_add_scalar_f32, float)
@@ -350,12 +350,12 @@ static inline void hvx_max_scalar_f32(uint8_t * restrict dst, const uint8_t * re
// CLAMP Scalar variants
#define HVX_OP_CLAMP_SCALAR(v) \
({ \
#define HVX_OP_CLAMP_SCALAR(v) \
({ \
HVX_VectorPred pred_cap_right = Q6_Q_vcmp_gt_VsfVsf(v, max_vec); \
HVX_VectorPred pred_cap_left = Q6_Q_vcmp_gt_VsfVsf(min_vec, v); \
HVX_Vector tmp = Q6_V_vmux_QVV(pred_cap_right, max_vec, v); \
Q6_V_vmux_QVV(pred_cap_left, min_vec, tmp); \
HVX_Vector tmp = Q6_V_vmux_QVV(pred_cap_right, max_vec, v); \
Q6_V_vmux_QVV(pred_cap_left, min_vec, tmp); \
})
static inline void hvx_clamp_scalar_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, const float min, const float max, uint32_t n) {
+47 -47
View File
@@ -219,64 +219,64 @@ static inline HVX_Vector hvx_vec_hybrid_div_f16(HVX_Vector vec1, HVX_Vector vec2
} while(0)
// Generic macro to define alignment permutations for an op
#define DEFINE_HVX_DIV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY) \
#define DEFINE_HVX_DIV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY) \
static inline void OP_NAME##_aaa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_Vector, hvx_vec_store_a); \
} \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_Vector, hvx_vec_store_a); \
} \
static inline void OP_NAME##_aau(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_UVector, hvx_vec_store_a); \
} \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src0 % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_Vector, HVX_UVector, hvx_vec_store_a); \
} \
static inline void OP_NAME##_aua(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_Vector, hvx_vec_store_a); \
} \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_Vector, hvx_vec_store_a); \
} \
static inline void OP_NAME##_auu(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_UVector, hvx_vec_store_a); \
} \
assert((uintptr_t) dst % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_UVector, HVX_UVector, hvx_vec_store_a); \
} \
static inline void OP_NAME##_uaa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_Vector, hvx_vec_store_u); \
} \
assert((uintptr_t) src0 % 128 == 0); \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_Vector, hvx_vec_store_u); \
} \
static inline void OP_NAME##_uau(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
assert((uintptr_t) src0 % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_UVector, hvx_vec_store_u); \
} \
assert((uintptr_t) src0 % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_Vector, HVX_UVector, hvx_vec_store_u); \
} \
static inline void OP_NAME##_uua(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_Vector, hvx_vec_store_u); \
} \
assert((uintptr_t) src1 % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_Vector, hvx_vec_store_u); \
} \
static inline void OP_NAME##_uuu(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) { \
OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_UVector, hvx_vec_store_u); \
} \
OP_LOOP_BODY(HVX_UVector, HVX_UVector, HVX_UVector, hvx_vec_store_u); \
} \
// Dispatcher logic
#define HVX_DIV_DISPATCHER(OP_NAME) \
#define HVX_DIV_DISPATCHER(OP_NAME) \
static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, const uint32_t num_elems) { \
if (hex_is_aligned((void *) dst, 128)) { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems); \
else OP_NAME##_aau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems); \
else OP_NAME##_auu(dst, src0, src1, num_elems); \
} \
} else { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems); \
else OP_NAME##_uau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems); \
else OP_NAME##_uuu(dst, src0, src1, num_elems); \
} \
} \
if (hex_is_aligned((void *) dst, 128)) { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aaa(dst, src0, src1, num_elems); \
else OP_NAME##_aau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_aua(dst, src0, src1, num_elems); \
else OP_NAME##_auu(dst, src0, src1, num_elems); \
} \
} else { \
if (hex_is_aligned((void *) src0, 128)) { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uaa(dst, src0, src1, num_elems); \
else OP_NAME##_uau(dst, src0, src1, num_elems); \
} else { \
if (hex_is_aligned((void *) src1, 128)) OP_NAME##_uua(dst, src0, src1, num_elems); \
else OP_NAME##_uuu(dst, src0, src1, num_elems); \
} \
} \
}
DEFINE_HVX_DIV_OP_VARIANTS(hvx_div_f32, hvx_div_f32_loop_body)
+23 -23
View File
@@ -169,36 +169,36 @@ static inline HVX_Vector hvx_vec_inverse_f16_guard(HVX_Vector v_sf, HVX_Vector n
} while(0)
// Generic macro to define alignment permutations for an op
#define DEFINE_HVX_INV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY) \
#define DEFINE_HVX_INV_OP_VARIANTS(OP_NAME, OP_LOOP_BODY) \
static inline void OP_NAME##_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_Vector, hvx_vec_store_a); \
} \
assert((uintptr_t) dst % 128 == 0); \
assert((uintptr_t) src % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_Vector, hvx_vec_store_a); \
} \
static inline void OP_NAME##_au(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
assert((uintptr_t) dst % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_UVector, hvx_vec_store_a); \
} \
assert((uintptr_t) dst % 128 == 0); \
OP_LOOP_BODY(HVX_Vector, HVX_UVector, hvx_vec_store_a); \
} \
static inline void OP_NAME##_ua(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
assert((uintptr_t) src % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_Vector, hvx_vec_store_u); \
} \
assert((uintptr_t) src % 128 == 0); \
OP_LOOP_BODY(HVX_UVector, HVX_Vector, hvx_vec_store_u); \
} \
static inline void OP_NAME##_uu(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) { \
OP_LOOP_BODY(HVX_UVector, HVX_UVector, hvx_vec_store_u); \
} \
OP_LOOP_BODY(HVX_UVector, HVX_UVector, hvx_vec_store_u); \
} \
// Dispatcher logic
#define HVX_INV_DISPATCHER(OP_NAME) \
#define HVX_INV_DISPATCHER(OP_NAME) \
static inline void OP_NAME(uint8_t * restrict dst, const uint8_t * restrict src, const uint32_t num_elems) { \
if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) { \
OP_NAME##_aa(dst, src, num_elems); \
} else if (hex_is_aligned((void *) dst, 128)) { \
OP_NAME##_au(dst, src, num_elems); \
} else if (hex_is_aligned((void *) src, 128)) { \
OP_NAME##_ua(dst, src, num_elems); \
} else { \
OP_NAME##_uu(dst, src, num_elems); \
} \
if (hex_is_aligned((void *) dst, 128) && hex_is_aligned((void *) src, 128)) { \
OP_NAME##_aa(dst, src, num_elems); \
} else if (hex_is_aligned((void *) dst, 128)) { \
OP_NAME##_au(dst, src, num_elems); \
} else if (hex_is_aligned((void *) src, 128)) { \
OP_NAME##_ua(dst, src, num_elems); \
} else { \
OP_NAME##_uu(dst, src, num_elems); \
} \
}
DEFINE_HVX_INV_OP_VARIANTS(hvx_inverse_f32, hvx_inverse_f32_loop_body)
+22 -22
View File
@@ -68,30 +68,30 @@ static inline void hvx_scale_f32(uint8_t * restrict dst, const uint8_t * restric
}
}
#define hvx_scale_offset_f32_loop_body(dst_type, src_type, vec_store) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
HVX_Vector vs = hvx_vec_splat_f32(scale); \
HVX_Vector vo = hvx_vec_splat_f32(offset); \
\
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) { \
#define hvx_scale_offset_f32_loop_body(dst_type, src_type, vec_store) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
HVX_Vector vs = hvx_vec_splat_f32(scale); \
HVX_Vector vo = hvx_vec_splat_f32(offset); \
\
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) { \
HVX_Vector v = Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(vsrc[i], vs), vo); \
vdst[i] = Q6_Vsf_equals_Vqf32(v); \
} \
if (nloe) { \
vdst[i] = Q6_Vsf_equals_Vqf32(v); \
} \
if (nloe) { \
HVX_Vector v = Q6_Vqf32_vadd_Vqf32Vsf(Q6_Vqf32_vmpy_VsfVsf(vsrc[i], vs), vo); \
vec_store((void *) &vdst[i], nloe * elem_size, Q6_Vsf_equals_Vqf32(v)); \
} \
vec_store((void *) &vdst[i], nloe * elem_size, Q6_Vsf_equals_Vqf32(v)); \
} \
} while(0)
static inline void hvx_scale_offset_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, const int n, const float scale, const float offset) {
+40 -40
View File
@@ -68,50 +68,50 @@ static inline HVX_Vector hvx_vec_tanh_f32(HVX_Vector x) {
return Q6_Vsf_equals_Vqf32(res);
}
#define hvx_sigmoid_loop_body(dst_type, src_type, vec_store) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
const HVX_Vector one = hvx_vec_splat_f32(1.f); \
const HVX_Vector max_exp = hvx_vec_splat_f32(87.f); \
const HVX_Vector min_exp = hvx_vec_splat_f32(-87.f); \
\
const uint32_t epv = 128 / sizeof(float); \
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_fast_sigmoid_f32_guard(vsrc[i], one, max_exp, min_exp); \
} \
if (nloe) { \
#define hvx_sigmoid_loop_body(dst_type, src_type, vec_store) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
const HVX_Vector one = hvx_vec_splat_f32(1.f); \
const HVX_Vector max_exp = hvx_vec_splat_f32(87.f); \
const HVX_Vector min_exp = hvx_vec_splat_f32(-87.f); \
\
const uint32_t epv = 128 / sizeof(float); \
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_fast_sigmoid_f32_guard(vsrc[i], one, max_exp, min_exp); \
} \
if (nloe) { \
HVX_Vector tmp = hvx_vec_fast_sigmoid_f32_guard(vsrc[i], one, max_exp, min_exp); \
vec_store((void *) &vdst[i], nloe * sizeof(float), tmp); \
} \
vec_store((void *) &vdst[i], nloe * sizeof(float), tmp); \
} \
} while(0)
#define hvx_tanh_loop_body(dst_type, src_type, vec_store) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
const uint32_t epv = 128 / sizeof(float); \
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_tanh_f32(vsrc[i]); \
} \
if (nloe) { \
HVX_Vector tmp = hvx_vec_tanh_f32(vsrc[i]); \
#define hvx_tanh_loop_body(dst_type, src_type, vec_store) \
do { \
dst_type * restrict vdst = (dst_type *) dst; \
src_type * restrict vsrc = (src_type *) src; \
\
const uint32_t epv = 128 / sizeof(float); \
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_tanh_f32(vsrc[i]); \
} \
if (nloe) { \
HVX_Vector tmp = hvx_vec_tanh_f32(vsrc[i]); \
vec_store((void *) &vdst[i], nloe * sizeof(float), tmp); \
} \
} \
} while(0)
static inline void hvx_sigmoid_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
+69 -36
View File
@@ -3,11 +3,12 @@
#pragma clang diagnostic ignored "-Wunused-but-set-variable"
#include <HAP_farf.h>
#include <HAP_perf.h>
#include <hexagon_protos.h>
#include <hexagon_types.h>
#include <string.h>
#include "hex-common.h"
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "htp-ctx.h"
@@ -16,14 +17,19 @@
#include "hex-dma.h"
#include "hex-profile.h"
#include "htp-vtcm.h"
#include "htp-tensor.h"
struct htp_im2col_context {
struct htp_ops_context * octx;
uint32_t patch_base; // first patch index assigned to this dev
uint32_t npatches; // number of patches assigned to this dev
uint32_t npatches_per_thread; // patches = N*OH*OW (pure-DDR kernel)
uint32_t pe_rows_per_thread; // N*OH rows per worker
uint32_t pe_src_row_bytes; // one output row's source: IC*KH*IW*4, rounded 256
uint32_t pe_dst_row_bytes; // one output row's dst: OW*patch_stride*2, rounded 256
uint32_t pe_row_base; // first N*OH row index assigned to this dev (DMA path)
uint32_t pe_nrows; // number of N*OH rows assigned to this dev (DMA path)
uint32_t pe_rows_per_thread; // N*OH rows per worker
uint32_t pe_src_row_bytes; // one output row's source: IC*KH*IW*4, rounded 256
uint32_t pe_dst_row_bytes; // one output row's dst: OW*patch_stride*2, rounded 256
// Patch-embed DMA path VTCM ping-pong.
uint8_t * pe_vtcm_src; // base of the 2x src buffers region
@@ -58,33 +64,27 @@ static inline void htp_im2col_vtcm_layout_build(struct htp_im2col_vtcm_layout *
struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \
struct htp_ops_context * octx = ictx->octx; \
struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \
const struct htp_tensor * restrict src0 = octx->src[0]; \
const struct htp_tensor * restrict src1 = octx->src[1]; \
const struct htp_tensor * restrict dst = octx->dst; \
const int32_t s0 = octx->op_params[0]; \
const int32_t s1 = octx->op_params[1]; \
const int32_t p0 = octx->op_params[2]; \
const int32_t p1 = octx->op_params[3]; \
const int32_t d0 = octx->op_params[4]; \
const int32_t d1 = octx->op_params[5]; \
const uint32_t N = src1->ne[3]; \
const uint32_t IC = src1->ne[2]; \
const uint32_t IH = src1->ne[1]; \
const uint32_t IW = src1->ne[0]; \
const uint32_t KH = octx->src[0]->ne[1]; \
const uint32_t KW = octx->src[0]->ne[0]; \
const int32_t s0 = octx->op_params[0], s1 = octx->op_params[1]; \
const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3]; \
const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5]; \
const uint32_t N = src1->ne[3], IC = src1->ne[2], IH = src1->ne[1], IW = src1->ne[0]; \
const uint32_t KH = src0->ne[1], KW = src0->ne[0]; \
const uint32_t OH = dst->ne[2]; \
const uint32_t OW = dst->ne[1]; \
const uint32_t patch_stride = IC * KH * KW; \
const float * restrict src_data = (const float *) src1->data; \
DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \
const uint32_t npatches = N * OH * OW; \
const uint32_t patch_start = ictx->npatches_per_thread * ith; \
const uint32_t patch_end = MIN(patch_start + ictx->npatches_per_thread, npatches); \
if (patch_start >= patch_end) { \
const uint32_t patch_end = ictx->patch_base + ictx->npatches; \
const uint32_t patch_start = ictx->patch_base + ictx->npatches_per_thread * ith; \
const uint32_t patch_stop = MIN(patch_start + ictx->npatches_per_thread, patch_end);\
if (patch_start >= patch_stop) { \
return; \
} \
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \
for (uint32_t p = patch_start; p < patch_end; p++) { \
for (uint32_t p = patch_start; p < patch_stop; p++) { \
const uint32_t iow = p % OW; \
const uint32_t ioh = (p / OW) % OH; \
const uint32_t in = p / (OW * OH); \
@@ -154,10 +154,10 @@ IM2COL_PATCHEMBED_BODY(im2col_patchembed_f32_thread, float, hvx_copy_f32_uu, hvx
uint8_t * dst_base = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread; \
float * srcb = (float *) src_base; \
DST_CTYPE * dstb = (DST_CTYPE *) dst_base; \
const uint32_t nrows = N * OH; \
const uint32_t row_end_max = ictx->pe_row_base + ictx->pe_nrows; \
const uint32_t per_thread = ictx->pe_rows_per_thread; \
const uint32_t row_start = per_thread * ith; \
const uint32_t row_end = MIN(row_start + per_thread, nrows); \
const uint32_t row_start = ictx->pe_row_base + per_thread * ith; \
const uint32_t row_end = MIN(row_start + per_thread, row_end_max); \
if (row_start >= row_end) \
return; \
for (uint32_t r = row_start; r < row_end; r++) { \
@@ -266,26 +266,55 @@ int op_im2col(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
const uint32_t N = src1->ne[3];
const uint32_t OH = dst->ne[2];
const uint32_t OW = dst->ne[1];
const uint32_t npatches = N * OH * OW;
const uint32_t n_threads = MIN(octx->n_threads, npatches);
if ((octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) || n_threads == 0) {
if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
return HTP_STATUS_OK;
}
const uint32_t N = src1->ne[3];
const uint32_t OH = dst->ne[2];
const uint32_t OW = dst->ne[1];
const uint32_t total_patches = N * OH * OW;
const uint32_t total_rows = N * OH;
uint32_t patch_base = 0;
uint32_t npatches = total_patches;
if (octx->ctx->mdev.count > 1) {
const uint32_t patch_size = dst->nb[1];
const uint32_t patches_per_chunk = (patch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(patch_size, HEX_L2_LINE_SIZE)) : 1;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_patches, htp_tensor_mdev_data_aligned(dst) ? patches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
patch_base = range.start;
npatches = range.count;
}
uint32_t row_base = 0;
uint32_t nrows = total_rows;
if (octx->ctx->mdev.count > 1) {
const uint32_t row_size = dst->nb[2];
const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
row_base = range.start;
nrows = range.count;
}
if (npatches == 0 && nrows == 0) {
return HTP_STATUS_OK;
}
const uint32_t n_threads = MIN(octx->n_threads, MAX(npatches, 1));
struct htp_im2col_context ictx = { 0 };
ictx.octx = octx;
ictx.npatches_per_thread = (npatches + n_threads - 1) / n_threads;
ictx.octx = octx;
ictx.patch_base = patch_base;
ictx.npatches = npatches;
ictx.npatches_per_thread = (npatches + n_threads - 1) / n_threads;
// Clean non-overlapping patch-embed -> DMA kernel (if it fits VTCM);
// everything else (padding/dilation/stride edges) -> pure-DDR kernel.
if (im2col_use_patchembed_dma(octx)) {
const uint32_t nrows = N * OH;
const uint32_t pth = MIN(octx->n_threads, nrows);
if (im2col_use_patchembed_dma(octx) && nrows > 0) {
const uint32_t pth = MIN(octx->n_threads, nrows);
if (pth > 0 && im2col_patchembed_dma_fits(octx, &ictx, pth)) {
ictx.pe_row_base = row_base;
ictx.pe_nrows = nrows;
ictx.pe_rows_per_thread = (nrows + pth - 1) / pth;
if (dst->type == HTP_TYPE_F16) {
work_queue_run(octx->ctx->work_queue, im2col_patchembed_dma_thread, &ictx, pth);
@@ -297,6 +326,10 @@ int op_im2col(struct htp_ops_context * octx) {
// else: doesn't fit -> fall through to the pure-DDR kernel below.
}
if (npatches == 0) {
return HTP_STATUS_OK;
}
if (dst->type == HTP_TYPE_F16) {
work_queue_run(octx->ctx->work_queue, im2col_patchembed_thread, &ictx, n_threads);
} else {
+80 -23
View File
@@ -34,6 +34,7 @@
#include "work-queue.h"
#include "hex-profile.h"
#include "allreduce-ops.h"
#include "htp-fence.h"
#define HMX_QUEUE_CAPACITY 16
#define HMX_QUEUE_STACK_SIZE 16384
@@ -710,22 +711,43 @@ static inline void profile_stop(uint32_t mode, struct profile_data * d) {
static int op_fence(struct htp_ops_context * octx) {
struct htp_context *ctx = octx->ctx;
struct htp_thread_trace * tr = &ctx->trace[0];
const uint32_t seq = (uint32_t) octx->op_params[0];
const uint32_t seq = (uint32_t) octx->op_params[0];
const uint32_t mode = (uint32_t) octx->op_params[1];
htp_trace_event_start(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
const struct htp_tensor * sync = octx->src[0];
atomic_uint * sync_fence = (atomic_uint *) sync->data;
atomic_uint * sync_fence = (atomic_uint *) (uintptr_t) sync->data;
if (mode == 1) {
htp_flush_dirty_ranges(ctx);
htp_mdev_group_barrier(octx);
if (ctx->mdev.idx == 0) {
htp_fence_write(sync_fence, seq, octx->status);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
FARF(HIGH, "ggml-hex: sync-signal : fence %p seq 0x%x status %d\n", sync_fence, seq, octx->status);
return octx->status;
}
int status = HTP_STATUS_OK;
uint64_t spins = 0;
while (1) {
Q6_dccleaninva_A((void *) sync_fence);
asm volatile ("syncht" : : : "memory");
uint32_t val = atomic_load(&sync_fence[0]);
if ((int32_t)(val - seq) >= 0) {
uint32_t sync_seq;
uint32_t sync_status;
htp_fence_read(sync_fence, &sync_seq, &sync_status);
if ((int32_t)(sync_seq - seq) >= 0) {
if (sync_status > HTP_STATUS_OK) {
FARF(ERROR, "ggml-hex: sync-wait peer failed with status %u : fence %p seq 0x%x\n", sync_status, sync_fence, seq);
status = sync_status;
}
break;
}
if (++spins > HTP_FENCE_TIMEOUT) {
FARF(ERROR, "ggml-hex: sync-wait TIMEOUT : fence %p spins %llu seq %u\n", sync_fence, spins, seq);
FARF(ERROR, "ggml-hex: sync-wait TIMEOUT : fence %p spins %llu seq 0x%x\n", sync_fence, spins, seq);
status = HTP_STATUS_INTERNAL_ERR;
break;
}
hex_pause();
@@ -733,12 +755,27 @@ static int op_fence(struct htp_ops_context * octx) {
htp_trace_event_stop(tr, HTP_TRACE_EVT_FENCE, (uint16_t) seq);
FARF(HIGH, "ggml-hex: sync-done : fence %p spins %llu seq %u\n", sync_fence, spins, seq);
FARF(HIGH, "ggml-hex: sync-done : fence %p spins %llu seq 0x%x\n", sync_fence, spins, seq);
return status;
}
static int op_mdev_group(struct htp_ops_context * octx) {
struct htp_context * ctx = octx->ctx;
const struct htp_tensor * sync = octx->src[0];
ctx->mdev.idx = (uint16_t) octx->op_params[0];
ctx->mdev.count = (uint16_t) sync->ne[1];
if (ctx->mdev.count > 1) {
ctx->mdev.count_div = init_fastdiv_values(ctx->mdev.count);
ctx->mdev.fence_base = (uint8_t *) sync->data;
}
return HTP_STATUS_OK;
}
static int execute_op(struct htp_ops_context * octx) {
switch (octx->op) {
case HTP_OP_MDEV_GROUP:
return op_mdev_group(octx);
case HTP_OP_FENCE:
return op_fence(octx);
@@ -812,6 +849,7 @@ static int execute_op(struct htp_ops_context * octx) {
return op_sum_rows(octx);
case HTP_OP_CPY:
case HTP_OP_CPY_FENCE:
return op_cpy(octx);
case HTP_OP_REPEAT:
@@ -855,7 +893,7 @@ static int execute_op(struct htp_ops_context * octx) {
}
FARF(ERROR, "Unknown Op %u", octx->op);
return -1;
return HTP_STATUS_NO_SUPPORT;
}
static inline bool reuse_buf(struct htp_context *ctx, uint32_t *m_reuse, struct htp_buf_desc *b) {
@@ -984,11 +1022,19 @@ static void prep_tensors(struct htp_context *ctx, struct htp_buf_desc *bufs, str
}
}
static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, uint32_t idx, struct htp_op_desc * op) {
memcpy(octx->op_params, op->params, sizeof(octx->op_params));
static void mdev_group_init(struct htp_context * ctx, const struct htp_opbatch_req * req) {
memset(&ctx->mdev, 0, sizeof(ctx->mdev));
ctx->mdev.fence_seq = (uint32_t)((req->seq & 0xfffff) << 12);
}
static int proc_op_req(struct htp_ops_context * octx, struct htp_buf_desc * bufs, uint32_t n_bufs,
struct htp_tensor * tens, uint32_t idx, struct htp_op_desc * op) {
memcpy(octx->op_params, op->params, sizeof(octx->op_params));
memcpy(octx->kernel_params, op->kernel_params, sizeof(octx->kernel_params));
octx->flags = op->flags;
octx->op = op->opcode;
octx->flags = op->flags;
octx->op = op->opcode;
octx->n_threads = octx->ctx->n_threads;
octx->n_threads_div = octx->ctx->n_threads_div;
FARF(HIGH, "proc-op #%u: opcode %u flags 0x%x", idx, octx->op, octx->flags);
@@ -1027,9 +1073,13 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, u
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]);
}
htp_tensor_dirty_all(octx->ctx, octx->dsts, HTP_OP_MAX_OUTPUTS);
htp_mdev_group_barrier(octx);
int status = execute_op(octx);
htp_tensor_dirty_all(octx->ctx, octx->dsts, HTP_OP_MAX_OUTPUTS);
htp_ops_context_set_status(octx, status);
octx->src0_spad.src = NULL;
octx->src1_spad.src = NULL;
@@ -1037,7 +1087,7 @@ static int proc_op_req(struct htp_ops_context * octx, struct htp_tensor *tens, u
octx->src3_spad.src = NULL;
octx->dst_spad.src = NULL;
return status;
return octx->status;
}
static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_req * req, const struct dspqueue_buffer * dbuf) {
@@ -1059,7 +1109,7 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r
return;
}
FARF(HIGH, "processing opbatch #%u: n-bufs %u n-tensors %u n-ops %u n-traces %u : m-size %u b-size %u t-size %u o-size %u", req->id,
FARF(HIGH, "processing opbatch #%llu: n-bufs %u n-tensors %u n-ops %u n-traces %u : m-size %u b-size %u t-size %u o-size %u", (unsigned long long) req->seq,
n_bufs, n_tens, n_ops, req->n_traces, dbuf->size, b_size, t_size, o_size);
// Setup descriptor pointers
@@ -1096,8 +1146,11 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r
struct htp_ops_context *octx = &ctx->octx;
memset(octx, 0, sizeof(*octx));
octx->n_threads = ctx->n_threads;
octx->ctx = ctx;
octx->n_threads = ctx->n_threads;
octx->n_threads_div = ctx->n_threads_div;
octx->ctx = ctx;
mdev_group_init(ctx, req);
work_queue_wakeup(ctx->work_queue);
if (ctx->hmx_queue) {
@@ -1105,15 +1158,18 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r
}
int op_status = HTP_STATUS_OK;
for (uint32_t i = 0; i < n_ops && op_status == HTP_STATUS_OK; i++) {
octx->status = HTP_STATUS_OK;
for (uint32_t i = 0; i < n_ops; i++) {
struct profile_data prof;
profile_start(ctx->profiler, &prof);
op_status = proc_op_req(octx, tens, i, &ops[i]);
op_status = proc_op_req(octx, bufs, n_bufs, tens, i, &ops[i]);
profile_stop(ctx->profiler, &prof);
htp_ops_context_set_status(octx, op_status);
if (ctx->profiler) {
pds[i].opcode = ops[i].opcode;
pds[i].usecs = prof.usecs;
@@ -1136,19 +1192,20 @@ static void process_opbatch(struct htp_context * ctx, const struct htp_opbatch_r
qurt_mem_cache_clean((qurt_addr_t) 0, 0, QURT_MEM_CACHE_FLUSH_INVALIDATE_ALL, QURT_MEM_DCACHE);
htp_trace_event_stop(&ctx->trace[0], HTP_TRACE_EVT_L2FLUSH, 0);
htp_mdev_group_barrier(octx);
profile_stop(HTP_PROF_BASIC, &batch_prof);
struct htp_opbatch_rsp rsp;
memset(&rsp, 0, sizeof(rsp));
rsp.id = req->id;
rsp.status = op_status;
rsp.seq = req->seq;
rsp.status = octx->status;
rsp.n_bufs = n_bufs;
rsp.n_tensors = n_tens;
rsp.n_ops = n_ops;
rsp.usecs = batch_prof.usecs;
rsp.cycles_start = batch_prof.cycles_start;
rsp.cycles_stop = batch_prof.cycles_stop;
rsp.seq = req->seq;
if (ctx->profiler == HTP_PROF_TRACE) {
for (int t = 0; t <= HTP_MAX_NTHREADS; t++) {
+270 -101
View File
@@ -21,6 +21,7 @@
#include "ggml-common.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#include "matmul-ops.h"
#include "htp-vtcm.h"
@@ -89,6 +90,8 @@ struct htp_mm_context {
// Precomputed values
uint32_t src0_nrows_per_thread;
uint32_t src0_row_start;
uint32_t src0_row_end;
uint32_t src0_row_size_padded;
uint32_t src1_nrows;
@@ -135,6 +138,23 @@ struct htp_mm_context {
uint32_t vtcm_dst_size_per_thread;
};
static int htp_mm_init_context(
struct htp_ops_context * octx,
const struct htp_mm_kernel_params * kparams
) {
if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
if (kparams->n_hmx) {
if (kparams->n_act_threads <= 0 || kparams->n_act_threads > (int32_t) octx->n_threads) {
return HTP_STATUS_INVAL_PARAMS;
}
}
return HTP_STATUS_OK;
}
// vdelta control to expand first 32 e8m0 values into 32 uint32 elements
static const uint8_t __attribute__((aligned(128))) expand_x32_e8m0[128] = {
0x00, 0x00, 0x00, 0x00, 0x01, 0x04, 0x00, 0x00, 0x02, 0x00, 0x08, 0x08, 0x01, 0x02, 0x00, 0x04, 0x04, 0x00, 0x00,
@@ -238,22 +258,24 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
// This is the size of the rest of the dimensions of the result
const uint32_t nr1 = ne1 * ne2 * ne3;
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;
// distribute the thread work across the inner or outer loop based on which one is larger
uint32_t dr0, dr1, ith0, ith1;
if (nr0 > nr1) {
dr0 = fastdiv(nr0 + nth - 1, &octx->ctx->n_threads_div);
dr0 = fastdiv(src0_nrows + nth - 1, &octx->n_threads_div);
dr1 = nr1;
ith0 = ith;
ith1 = 0;
} else {
dr0 = nr0;
dr1 = fastdiv(nr1 + nth - 1, &octx->ctx->n_threads_div);
dr0 = src0_nrows;
dr1 = fastdiv(nr1 + nth - 1, &octx->n_threads_div);
ith0 = 0;
ith1 = ith;
}
const uint32_t ir0_start = dr0 * ith0;
const uint32_t ir0_end = MIN(ir0_start + dr0, nr0);
const uint32_t ir0_start = mmctx->src0_row_start + dr0 * ith0;
const uint32_t ir0_end = MIN(ir0_start + dr0, mmctx->src0_row_end);
const uint32_t ir1_start = dr1 * ith1;
const uint32_t ir1_end = MIN(ir1_start + dr1, nr1);
@@ -312,11 +334,11 @@ static void hvx_mm_4d(unsigned int nth, unsigned int ith, void * data) {
static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \
htp_matmul_preamble; \
\
const uint32_t src0_nrows = ne01 * ne02 * ne03; \
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; \
const uint32_t src1_nrows = ne11 * ne12 * ne13; \
\
const uint32_t src0_start_row = src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); \
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end); \
\
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
\
@@ -414,10 +436,10 @@ static void hvx_mm_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void * data) { \
htp_matmul_preamble; \
\
const uint32_t src0_nrows = ne01; \
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; \
\
const uint32_t src0_start_row = src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); \
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end); \
\
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
\
@@ -549,12 +571,22 @@ static void hvx_mm_nx_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, v
uint32_t n_k_tiles_w = ne00 / 32; \
uint32_t tile_row_stride = n_k_tiles_w * tile_size; \
\
const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3]; \
uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div); \
uint32_t src0_start_row = 0; \
uint32_t src0_end_row = ne01; \
if (octx->ctx->mdev.count > 1) { \
const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float)); \
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0, \
octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div); \
src0_start_row = range.start; \
src0_end_row = range.start + range.count; \
} \
\
const uint32_t nrows = src0_end_row - src0_start_row; \
uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div); \
src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32); \
\
const uint32_t start_row = src0_nrows_per_thread * ith; \
const uint32_t end_row = MIN(start_row + src0_nrows_per_thread, src0_nrows); \
const uint32_t start_row = src0_start_row + src0_nrows_per_thread * ith; \
const uint32_t end_row = MIN(start_row + src0_nrows_per_thread, src0_end_row); \
if (start_row >= end_row) continue; \
\
uint32_t ct_start = start_row / 32; \
@@ -735,11 +767,11 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
assert(n_prefetch >= 2 && n_prefetch <= HTP_MM_MAX_PREFETCH && (n_prefetch & (n_prefetch - 1)) == 0);
const uint32_t prefetch_mask = n_prefetch - 1;
const uint32_t src0_nrows = ne01 * ne02 * ne03; // src0 rows
const uint32_t src1_nrows = ne11 * ne12 * ne13; // src1 rows
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; // src0 rows
const uint32_t src1_nrows = ne11 * ne12 * ne13; // src1 rows
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U);
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
@@ -781,7 +813,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
const uint8_t * ss0 = dma_queue_pop(dma_queue).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ir0);
// Process src1 columns in pairs (2×2 tiling)
// Process src1 columns in pairs (2x2 tiling)
uint32_t ir1 = 0;
for (; ir1 + 1 < src1_nrows; ir1 += 2) {
const uint8_t * restrict src1_col0 = (const uint8_t *) (src1_data + (ir1+0) * src1_stride);
@@ -791,7 +823,7 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
mmctx->vec_dot_2x2(ne00, &dst_row0[ir0], &dst_row1[ir0], ss0, ss0 + src0_stride, src1_col0, src1_col1);
}
// Handle remaining src1 rows (fallback to 2×1)
// Handle remaining src1 rows (fallback to 2x1)
for (; ir1 < src1_nrows; ++ir1) {
const uint8_t * restrict src1_col = (const uint8_t *) (src1_data + ir1 * src1_stride);
float * restrict dst_row = (float *) (dst->data + (ir1 * dst_row_size));
@@ -833,10 +865,10 @@ static void hvx_mm_2d(unsigned int nth, unsigned int ith, void * data) {
static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
htp_matmul_preamble;
const uint32_t src0_nrows = ne01;
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start;
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
@@ -943,13 +975,10 @@ static void hvx_mm_id(unsigned int nth, unsigned int ith, void * data) {
const struct htp_tensor * restrict ids = octx->src[2];
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
const uint32_t src0_nrows = ne01; // src0 rows per expert
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; // src0 rows per expert
const uint32_t src1_nrows = ne11;
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
hvx_mm_run_quant_task(mmctx, ith);
@@ -1036,9 +1065,9 @@ static void hvx_mv_id(unsigned int nth, unsigned int ith, void * data) {
const struct htp_tensor * restrict ids = octx->src[2];
const uint32_t src0_nrows = ne01; // src0 rows per expert
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_nrows = mmctx->src0_row_end - mmctx->src0_row_start; // src0 rows per expert
const uint32_t src0_start_row = mmctx->src0_row_start + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, mmctx->src0_row_end);
hvx_mm_run_quant_task(mmctx, ith);
@@ -1143,12 +1172,22 @@ static void hvx_mv_id_nx(unsigned int nth, unsigned int ith, void * data) {
const struct htp_tensor * restrict dst = octx->dsts[p];
if (!src_w || !dst) continue;
const uint32_t src0_nrows = src_w->ne[1];
uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div);
const uint32_t ne01 = src_w->ne[1];
uint32_t start_row = 0;
uint32_t end_row = ne01;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
start_row = range.start;
end_row = range.start + range.count;
}
const uint32_t nrows = end_row - start_row;
uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div);
src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32);
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = start_row + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, end_row);
if (src0_start_row >= src0_end_row) continue;
const uint8_t * restrict src0_row = (const uint8_t *) src_w->data + eid * src_w->nb[2];
@@ -1227,12 +1266,22 @@ static void hvx_mm_id_nx(unsigned int nth, unsigned int ith, void * data) {
const struct htp_tensor * restrict dst = octx->dsts[p];
if (!src_w || !dst) continue;
const uint32_t src0_nrows = src_w->ne[1];
uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div);
const uint32_t ne01 = src_w->ne[1];
uint32_t start_row = 0;
uint32_t end_row = ne01;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
start_row = range.start;
end_row = range.start + range.count;
}
const uint32_t nrows = end_row - start_row;
uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div);
src0_nrows_per_thread = hex_round_up(src0_nrows_per_thread, 32);
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = start_row + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, end_row);
if (src0_start_row >= src0_end_row) continue;
const uint8_t * src0_row = (const uint8_t *) src_w->data + cur_a * src_w->nb[2];
@@ -1323,15 +1372,33 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
const uint32_t src0_nrows = ne01 * ne02 * ne03;
const uint32_t src0_nrows = ne01;
const uint32_t src1_nrows = ne11 * ne12 * ne13;
uint32_t src0_row_start = 0;
uint32_t src0_row_end = src0_nrows;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
src0_row_start = range.start;
src0_row_end = range.start + range.count;
}
if (src0_row_start >= src0_row_end) {
return HTP_STATUS_OK;
}
const uint32_t nrows = src0_row_end - src0_row_start;
mmctx->src0_row_start = src0_row_start;
mmctx->src0_row_end = src0_row_end;
bool is_repacked = (src0->type == HTP_TYPE_Q4_0 || src0->type == HTP_TYPE_Q4_1 ||
src0->type == HTP_TYPE_Q8_0 || src0->type == HTP_TYPE_IQ4_NL ||
src0->type == HTP_TYPE_MXFP4);
// Compute src0_nrows_per_thread
mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div);
mmctx->src0_nrows_per_thread = fastdiv(nrows + octx->n_threads - 1, &octx->n_threads_div);
if (is_repacked) {
mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
} else {
@@ -1503,13 +1570,13 @@ static int hvx_mm_matmul(struct htp_ops_context * octx) {
kparams->kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
} else {
mmctx->vtcm_src1_size_per_thread = fastdiv(L.src1_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_src1_size_per_thread = fastdiv(L.src1_bytes, &octx->n_threads_div);
}
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->n_threads_div);
size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
const size_t vtcm_size = L.total_bytes;
FARF(HIGH, "matmul-%s : src0-vtcm-size %zu src1-vtcm-size %zu dst-vtcm-size %zu (%zu)\n", mmctx->type,
L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
@@ -1583,13 +1650,21 @@ static void hvx_mm_nx_2d(unsigned int nth, unsigned int ith, void * data) {
const uint32_t ne00 = src_w->ne[0];
const uint32_t ne01 = src_w->ne[1];
const uint32_t src0_nrows = ne01 * src_w->ne[2] * src_w->ne[3];
uint32_t start_row = 0;
uint32_t end_row = ne01;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(ne01, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
start_row = range.start;
end_row = range.start + range.count;
}
uint32_t src0_nrows_per_thread = fastdiv(src0_nrows + nth - 1, &octx->ctx->n_threads_div);
const uint32_t nrows = end_row - start_row;
uint32_t src0_nrows_per_thread = fastdiv(nrows + nth - 1, &octx->n_threads_div);
src0_nrows_per_thread += (src0_nrows_per_thread & 1);
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = start_row + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, end_row);
const uint32_t src0_end_row_x2 = src0_start_row + ((src0_end_row - src0_start_row) & ~1U);
if (src0_start_row >= src0_end_row) continue;
@@ -2638,10 +2713,6 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
const struct htp_tensor * restrict src0 = octx->src[0];
const struct htp_tensor * restrict act = octx->src[n_weights];
if (!src0 || !act) {
return HTP_STATUS_INVAL_PARAMS;
}
const int weight_type = (int) src0->type;
const int k = (int) act->ne[0];
const int k_valid = (int) act->ne[0];
@@ -2714,16 +2785,31 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
hmx_init_column_scales(vtcm_scales, Q6_V_vsplat_R(0x3c00)); // scale: 1.0, bias: 0.0 in FP16
FARF(HIGH, "hmx-mm-nx-2d: n_weights %u m %d k %d wtype %d mc %d nc %d vtcm %zu/%zu",
n_weights, m, k, weight_type, m_chunk_n_rows, n_chunk_n_cols, L.total_bytes, vtcm_budget);
int m_start = 0;
int m_rows = m;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_can_row_partition(octx->dsts[0], sizeof(float));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_rows = (int) range.count;
}
if (m_rows == 0) {
return HTP_STATUS_OK;
}
FARF(HIGH, "hmx-mm-nx-2d: n_weights %u m %d (%d..%d) k %d wtype %d mc %d nc %d vtcm %zu/%zu",
n_weights, m, m_start, m_start + m_rows, k, weight_type, m_chunk_n_rows, n_chunk_n_cols, L.total_bytes, vtcm_budget);
htp_trace_event_stop(tr, HTP_TRACE_EVT_INIT, 0);
const size_t mr_end = (size_t)(m_start + m_rows);
if (pipeline) {
hmx_matmul_job_t job_slots[2];
for (size_t mr = 0; mr < (size_t) m; mr += m_chunk_n_rows) {
const size_t n_rows = hex_smin(m - mr, m_chunk_n_rows);
for (size_t mr = (size_t) m_start; mr < mr_end; mr += m_chunk_n_rows) {
const size_t n_rows = hex_smin(mr_end - mr, m_chunk_n_rows);
void *vtcm_weight_bufs[2] = { vtcm_scratch0, vtcm_scratch1 };
void *vtcm_output_bufs[2] = { vtcm_output, vtcm_scratch2 };
@@ -2822,8 +2908,8 @@ static int hmx_mm_nx_2d_f32(struct htp_ops_context * octx, const struct htp_mm_k
}
} else {
hmx_matmul_job_t job;
for (size_t mr = 0; mr < (size_t) m; mr += m_chunk_n_rows) {
const size_t n_rows = hex_smin(m - mr, m_chunk_n_rows);
for (size_t mr = (size_t) m_start; mr < mr_end; mr += m_chunk_n_rows) {
const size_t n_rows = hex_smin(mr_end - mr, m_chunk_n_rows);
struct activation_transfer_params act_params = {
.ctx = ctx,
@@ -3095,7 +3181,7 @@ static int hmx_mm_f16_f32_batched(struct htp_context *ctx, const hmx_mm_f16_f32_
int chunk_dst_cols = params->n - (int)nc;
if (chunk_dst_cols > 0) {
transfer_output_chunk_threaded(ctx, output, src2_chunk, vtcm_output, (int) n_rows, (int) n_cols,
params->dst_stride, params->src2_stride, chunk_dst_cols, ctx->n_threads);
params->dst_stride, params->src2_stride, chunk_dst_cols, n_threads);
}
}
}
@@ -3216,7 +3302,10 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
int weight_type,
const struct mmid_row_mapping *matrix_rows,
int cur_a,
int mapping_stride) {
int mapping_stride,
int m_start,
int m_end,
int n_threads) {
struct htp_thread_trace * tr = &ctx->trace[0];
htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
@@ -3247,7 +3336,6 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
const int n_k_tiles = k / HTP_MM_HMX_TILE_N_COLS;
const struct fastdiv_values n_k_tiles_div = init_fastdiv_values(n_k_tiles);
const int n_threads = ctx->n_threads;
const bool is_quant = (weight_type != HTP_TYPE_F16 && weight_type != HTP_TYPE_F32);
const size_t vec_dot_size = k * sizeof(__fp16);
@@ -3303,8 +3391,8 @@ static int hmx_mm_id_2d_f32(struct htp_context *ctx,
hmx_matmul_job_t job;
for (size_t mr = 0; mr < (size_t) m_padded; mr += m_chunk_n_rows) {
const size_t n_rows = hex_smin(m_padded - mr, m_chunk_n_rows);
for (size_t mr = (size_t) m_start; mr < (size_t) m_end; mr += m_chunk_n_rows) {
const size_t n_rows = hex_smin((size_t) m_end - mr, m_chunk_n_rows);
const size_t n_row_tiles = hmx_ceil_div(n_rows, HTP_MM_HMX_TILE_N_ROWS);
transfer_activation_chunk_gathered_threaded(
@@ -3368,31 +3456,48 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
const int act_stride = (int)(src1->nb[1] / sizeof(float));
const int wgt_stride = (int)(src0->nb[1] / sizeof(__fp16));
int m_start = 0;
int m_rows = m_total;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_total, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_rows = (int) range.count;
}
if (m_rows == 0) {
return HTP_STATUS_OK;
}
const float * src2_ptr = NULL;
uint32_t src2_stride = 0;
size_t src2_nb2 = 0;
size_t src2_nb3 = 0;
if (src2) {
src2_ptr = (const float *) src2->data;
src2_stride = (src2->ne[1] == 1) ? 0 : (uint32_t) (src2->nb[1] / sizeof(float));
src2_ptr = (const float *) src2->data + m_start * src2_stride;
src2_nb2 = (src2->ne[2] == 1) ? 0 : src2->nb[2];
src2_nb3 = (src2->ne[3] == 1) ? 0 : src2->nb[3];
}
const int dst_stride = (int)(dst->nb[1] / sizeof(float));
float * dst_ptr = (float *) dst->data + m_start * dst_stride;
const float * act_ptr = (const float *) src1->data + m_start * act_stride;
int ret = -1;
const int n_threads = MIN(kparams->n_threads, (int) octx->n_threads);
const int n_threads = kparams->n_threads;
if (kparams->kernel_type == HTP_MM_KERNEL_HMX_F16_BATCHED) {
hmx_mm_f16_f32_batched_params_t batch_params = {
.dst = (float *) dst->data,
.dst = dst_ptr,
.src2 = src2_ptr,
.activation = (float *) src1->data,
.activation = act_ptr,
.weight = (const __fp16 *) src0->data,
.m = m_total,
.m = m_rows,
.k = k,
.n = n,
.act_stride = act_stride,
.weight_stride = wgt_stride,
.dst_stride = (int) (dst->nb[1] / sizeof(float)),
.dst_stride = dst_stride,
.src2_stride = src2_stride,
.ne02 = ne02,
.ne03 = ne03,
@@ -3420,9 +3525,9 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
kparams->vtcm_size);
} else {
ret = hmx_mm_2d_f32(
octx->ctx, (float*) dst->data, src2_ptr, (float*) src1->data, (const uint8_t *) src0->data,
m_total, k, n, act_stride, (int) src0->nb[1], (int) src0->type, (int) src1->ne[0],
(int)(dst->nb[1] / sizeof(float)), src2_stride, (int)dst->ne[0],
octx->ctx, dst_ptr, src2_ptr, act_ptr, (const uint8_t *) src0->data,
m_rows, k, n, act_stride, (int) src0->nb[1], (int) src0->type, (int) src1->ne[0],
dst_stride, src2_stride, (int)dst->ne[0],
kparams->m_chunk, kparams->n_chunk, kparams->pipeline, n_threads,
kparams->n_act_threads,
&kparams->div_n_act_threads,
@@ -3441,6 +3546,11 @@ static int hmx_mm_op_matmul(struct htp_ops_context * octx, const struct htp_mm_k
int op_matmul(struct htp_ops_context * octx) {
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
const int status = htp_mm_init_context(octx, kparams);
if (status != HTP_STATUS_OK) {
return status;
}
if (kparams->n_hmx) {
return hmx_mm_op_matmul(octx, kparams);
}
@@ -3463,6 +3573,16 @@ static int hmx_mm_op_matmul_id(
const int32_t cne1 = matrix_row_counts[cur_a];
if (cne1 == 0) continue;
const int m_padded = hex_align_up(cne1, 32);
int m_start = 0, m_end = m_padded;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) && (uint32_t) cne1 >= octx->ctx->mdev.count;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_end = (int) (range.start + range.count);
}
if (m_start >= m_end) continue;
int ret = hmx_mm_id_2d_f32(octx->ctx, (float*) dst->data, (float*) src1->data,
(const uint8_t *) src0->data + cur_a * nb02,
cne1, ne00, ne01,
@@ -3471,7 +3591,8 @@ static int hmx_mm_op_matmul_id(
nb11, nb12,
nb1, nb2,
(int) src0->nb[1], (int) src0->type,
matrix_rows, cur_a, mmctx->mapping_stride);
matrix_rows, cur_a, mmctx->mapping_stride,
m_start, m_end, (int) octx->n_threads);
if (ret != 0) {
FARF(ERROR, "HMX matmul failed for expert %u, error %d\n", cur_a, ret);
return HTP_STATUS_NO_SUPPORT;
@@ -3524,7 +3645,7 @@ static int hvx_mm_matmul_id(
htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, ne10, src1_nrows, octx->n_threads,
0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);
size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
const size_t vtcm_size = L.total_bytes;
FARF(HIGH, "matmul-id-%s : src0-spad-size %zu src1-spad-size %zu src2-spad-size 0 dst-spad-size %zu (%zu)\n", mmctx->type,
L.src0_bytes, L.src1_bytes, L.dst_bytes, vtcm_size);
@@ -3554,10 +3675,10 @@ static int hvx_mm_matmul_id(
mmctx->vtcm_src0_stride = src0_row_size_padded;
mmctx->vtcm_src1_stride = src1_row_size;
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
mmctx->vtcm_src2_size_per_thread = 0;
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->n_threads_div);
mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
mmctx->quant_task_func = quant_task_func;
@@ -3587,6 +3708,20 @@ static int hmx_mm_op_matmul_id_nx(
const int32_t cne1 = matrix_row_counts[cur_a];
if (cne1 == 0) continue;
const int m_padded = hex_align_up(cne1, 32);
int m_start = 0, m_end = m_padded;
if (octx->ctx->mdev.count > 1) {
bool can_split = (uint32_t) cne1 >= octx->ctx->mdev.count;
for (uint32_t p = 0; p < n_weights && can_split; ++p) {
const struct htp_tensor * restrict dst = octx->dsts[p];
can_split = !dst || htp_tensor_mdev_data_aligned(dst);
}
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition((uint32_t) m_padded, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
m_start = (int) range.start;
m_end = (int) (range.start + range.count);
}
if (m_start >= m_end) continue;
for (uint32_t p = 0; p < n_weights; ++p) {
const struct htp_tensor * restrict src_w = octx->src[p];
const struct htp_tensor * restrict dst = octx->dsts[p];
@@ -3600,7 +3735,8 @@ static int hmx_mm_op_matmul_id_nx(
act->nb[1], act->nb[2],
dst->nb[1], dst->nb[2],
(int) src_w->nb[1], (int) src_w->type,
matrix_rows, cur_a, mmctx->mapping_stride);
matrix_rows, cur_a, mmctx->mapping_stride,
m_start, m_end, (int) octx->n_threads);
if (ret != 0) {
FARF(ERROR, "HMX matmul ID NX failed for expert %u weight %u, error %d\n", cur_a, p, ret);
return HTP_STATUS_NO_SUPPORT;
@@ -3656,7 +3792,7 @@ static int hvx_mm_matmul_id_nx(
htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads,
0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, true, false);
size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
const size_t vtcm_size = L.total_bytes;
if (octx->ctx->vtcm_size < vtcm_size) {
FARF(ERROR, "matmul-id-nx: current VTCM reservation %zu is too small, needed %zu\n",
@@ -3678,9 +3814,9 @@ static int hvx_mm_matmul_id_nx(
mmctx->vtcm_src0_stride = 0;
mmctx->vtcm_src1_stride = src1_row_size;
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->n_threads_div);
mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
mmctx->quant_task_func = quant_task_func;
@@ -3769,16 +3905,21 @@ static inline void scan_expert_ids(
int op_matmul_id(struct htp_ops_context * octx) {
htp_matmul_tensors_preamble;
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
struct htp_mm_context mmctx_struct = {0};
struct htp_mm_context * mmctx = &mmctx_struct;
const int status = htp_mm_init_context(octx, kparams);
if (status != HTP_STATUS_OK) {
return status;
}
struct htp_thread_trace * tr = &octx->ctx->trace[0];
htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
struct htp_mm_context mmctx_struct = {0};
struct htp_mm_context * mmctx = &mmctx_struct;
mmctx->octx = octx;
mmctx->act = src1;
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
const struct htp_tensor * restrict ids = octx->src[2];
const size_t src0_row_size = nb01;
@@ -3789,9 +3930,6 @@ int op_matmul_id(struct htp_ops_context * octx) {
const uint32_t src0_nrows = ne01; // per expert
const uint32_t src1_nrows = ne11 * ne12 * ne13;
mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div);
mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
// row groups
const int n_ids = ids->ne[0]; // n_expert_used
const int n_as = ne02; // n_expert
@@ -3843,6 +3981,29 @@ int op_matmul_id(struct htp_ops_context * octx) {
if (kparams->n_hmx) {
s = hmx_mm_op_matmul_id(octx, mmctx);
} else {
uint32_t src0_row_start = 0;
uint32_t src0_row_end = src0_nrows;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_can_row_partition(dst, sizeof(float));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, can_split ? 32 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
src0_row_start = range.start;
src0_row_end = range.start + range.count;
}
if (src0_row_start >= src0_row_end) {
if (mapping_buf != octx->ctx->ddr_spad_base) {
free(mapping_buf);
}
return HTP_STATUS_OK;
}
const uint32_t nrows = src0_row_end - src0_row_start;
mmctx->src0_row_start = src0_row_start;
mmctx->src0_row_end = src0_row_end;
mmctx->src0_nrows_per_thread = fastdiv(nrows + octx->n_threads - 1, &octx->n_threads_div);
mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
if (hvx_mm_init_vec_dot(mmctx, src0->type) == 0) {
s = hvx_mm_matmul_id(octx, mmctx, src1_nrows > 1 ? hvx_mm_id : hvx_mv_id);
} else {
@@ -3858,29 +4019,31 @@ int op_matmul_id(struct htp_ops_context * octx) {
}
int op_matmul_id_nx(struct htp_ops_context * octx) {
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
struct htp_mm_context mmctx_struct = {0};
struct htp_mm_context * mmctx = &mmctx_struct;
const int status = htp_mm_init_context(octx, kparams);
if (status != HTP_STATUS_OK) {
return status;
}
struct htp_thread_trace * tr = &octx->ctx->trace[0];
htp_trace_event_start(tr, HTP_TRACE_EVT_INIT, 0);
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
mmctx->octx = octx;
const uint32_t n_weights = kparams->n_weights;
const struct htp_tensor * restrict src0 = octx->src[0];
const struct htp_tensor * restrict act = octx->src[n_weights];
const struct htp_tensor * restrict ids = octx->src[n_weights + 1];
struct htp_mm_context mmctx_struct = {0};
struct htp_mm_context * mmctx = &mmctx_struct;
mmctx->octx = octx;
mmctx->act = act;
const size_t src0_row_size = src0->nb[1];
const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
const uint32_t src0_nrows = src0->ne[1];
const uint32_t src1_nrows = act->ne[1] * act->ne[2] * act->ne[3];
mmctx->src0_nrows_per_thread = fastdiv(src0_nrows + octx->n_threads - 1, &octx->ctx->n_threads_div);
mmctx->src0_nrows_per_thread = hex_round_up(mmctx->src0_nrows_per_thread, 32);
const int n_ids = ids->ne[0];
const int n_as = src0->ne[2];
@@ -3946,6 +4109,12 @@ int op_matmul_id_nx(struct htp_ops_context * octx) {
}
int op_matmul_nx(struct htp_ops_context * octx) {
const struct htp_mm_kernel_params * kparams = (const struct htp_mm_kernel_params *) octx->kernel_params;
const int status = htp_mm_init_context(octx, kparams);
if (status != HTP_STATUS_OK) {
return status;
}
if (kparams->n_hmx) {
return hmx_mm_nx_2d_f32(octx, kparams);
}
@@ -4012,7 +4181,7 @@ int op_matmul_nx(struct htp_ops_context * octx) {
htp_mm_hvx_vtcm_layout_build(&L, kparams->kernel_type, src0->type, act->ne[0], src1_nrows, octx->n_threads,
0, src0_row_size, src1_row_size, 0, kparams->n_prefetch, false, true);
size_t vtcm_size = kparams->vtcm_size > 0 ? (size_t)kparams->vtcm_size : L.total_bytes;
const size_t vtcm_size = L.total_bytes;
if (octx->ctx->vtcm_size < vtcm_size) {
FARF(ERROR, "matmul-nx: current VTCM reservation %zu is too small, needed %zu\n",
@@ -4034,9 +4203,9 @@ int op_matmul_nx(struct htp_ops_context * octx) {
mmctx->vtcm_src0_stride = is_repacked ? 0 : src0_row_size_padded;
mmctx->vtcm_src1_stride = src1_row_size;
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_src0_size_per_thread = fastdiv(L.src0_bytes, &octx->n_threads_div);
mmctx->vtcm_src1_size_per_thread = L.src1_bytes;
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->ctx->n_threads_div);
mmctx->vtcm_dst_size_per_thread = fastdiv(L.dst_bytes, &octx->n_threads_div);
mmctx->n_quant_rows_per_thread = (src1_nrows + n_quant_tasks - 1) / n_quant_tasks;
mmctx->quant_task_func = quant_task_func;
+85 -74
View File
@@ -12,8 +12,11 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
/* Circular wrap: maps any integer x into [0, n) */
static inline uint32_t wrap_around(int32_t x, uint32_t n) {
@@ -68,6 +71,7 @@ struct htp_pad_context {
uint32_t nrows_per_thread;
uint32_t total_dst_rows;
uint32_t row_start;
size_t type_size;
@@ -78,39 +82,39 @@ struct htp_pad_context {
size_t dst_row_size_aligned;
};
#define htp_pad_preamble \
const struct htp_tensor * src = octx->src[0]; \
const struct htp_tensor * dst = octx->dst; \
\
const uint32_t ne00 = src->ne[0]; \
const uint32_t nb00 = src->nb[0]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3]; \
\
const int32_t lp0 = pctx->lp0, rp0 = pctx->rp0; \
const int32_t lp1 = pctx->lp1, rp1 = pctx->rp1; \
const int32_t lp2 = pctx->lp2, rp2 = pctx->rp2; \
const int32_t lp3 = pctx->lp3, rp3 = pctx->rp3; \
\
const size_t type_size = pctx->type_size; \
\
const uint32_t row_start = pctx->nrows_per_thread * ith; \
const uint32_t row_end = MIN(row_start + pctx->nrows_per_thread, pctx->total_dst_rows);
#define htp_pad_preamble \
const struct htp_tensor * src = octx->src[0]; \
const struct htp_tensor * dst = octx->dst; \
\
const uint32_t ne00 = src->ne[0]; \
const uint32_t nb00 = src->nb[0]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3]; \
\
const int32_t lp0 = pctx->lp0, rp0 = pctx->rp0; \
const int32_t lp1 = pctx->lp1, rp1 = pctx->rp1; \
const int32_t lp2 = pctx->lp2, rp2 = pctx->rp2; \
const int32_t lp3 = pctx->lp3, rp3 = pctx->rp3; \
\
const size_t type_size = pctx->type_size; \
\
const uint32_t row_start = pctx->row_start + pctx->nrows_per_thread * ith; \
const uint32_t row_end = MIN(row_start + pctx->nrows_per_thread, pctx->row_start + pctx->total_dst_rows);
#define htp_pad_dma_preamble \
const size_t src_row_size = pctx->src_row_size; \
const size_t src_row_size_aligned = pctx->src_row_size_aligned; \
const size_t dst_row_size = pctx->dst_row_size; \
const size_t dst_row_size_aligned = pctx->dst_row_size_aligned; \
\
#define htp_pad_dma_preamble \
const size_t src_row_size = pctx->src_row_size; \
const size_t src_row_size_aligned = pctx->src_row_size_aligned; \
const size_t dst_row_size = pctx->dst_row_size; \
const size_t dst_row_size_aligned = pctx->dst_row_size_aligned; \
\
uint8_t * src_spad_base = octx->src0_spad.data + ith * octx->src0_spad.size_per_thread; \
uint8_t * dst_spad_base = octx->dst_spad.data + ith * octx->dst_spad.size_per_thread; \
\
@@ -125,8 +129,8 @@ static void pad_job_per_thread_hvx(unsigned int nth, unsigned int ith, void * da
struct htp_ops_context * octx = pctx->octx;
htp_pad_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
for (uint32_t dst_row = row_start; dst_row < row_end; dst_row++) {
uint32_t i1, i2, i3;
@@ -165,18 +169,17 @@ static void pad_job_per_thread_hvx(unsigned int nth, unsigned int ith, void * da
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
FARF(HIGH, "pad-hvx %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
FARF(HIGH, "pad-hvx %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
ith, nth,
src->ne[0], src->ne[1], src->ne[2], src->ne[3],
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
row_start, row_end,
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
row_start, row_end);
}
// ---------------------------------------------------------------------------
// HVX + DMA PAD kernel aligned, double-buffered
// HVX + DMA PAD kernel - aligned, double-buffered
// ---------------------------------------------------------------------------
static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void * data) {
@@ -185,9 +188,6 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
htp_pad_preamble;
htp_pad_dma_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
// -----------------------------------------------------------------------
// Priming phase: push 2 pairs of (dummy_dst_DMA, src_DMA) to seed the
// double-buffer pipeline before the main loop begins.
@@ -222,6 +222,8 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
// Main loop: pop completed DMAs, compute in VTCM with aligned HVX ops,
// push dst DMA and prefetch src for the next+1 row.
// -----------------------------------------------------------------------
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = row_start; ir < row_end; ir++) {
uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma).src;
uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma).dst;
@@ -236,6 +238,7 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
lp2, rp2, ne2,
lp3, rp3, ne3);
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
if (!interior) {
hvx_splat_f32_a(dst_spad_cur, 0.0f, ne0);
} else {
@@ -249,6 +252,7 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
hvx_copy_f32_ua(dst_interior, src_spad_cur, ne00);
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
dma_queue_push_vtcm_to_ddr(dma,
dma_make_ptr(dst_ptr, dst_spad_cur),
@@ -274,14 +278,11 @@ static void pad_job_per_thread_hvx_dma(unsigned int nth, unsigned int ith, void
dma_queue_flush(dma);
t2 = HAP_perf_get_qtimer_count();
FARF(HIGH, "pad-hvx-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
FARF(HIGH, "pad-hvx-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
ith, nth,
src->ne[0], src->ne[1], src->ne[2], src->ne[3],
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
row_start, row_end,
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
row_start, row_end);
}
// ---------------------------------------------------------------------------
@@ -293,8 +294,8 @@ static void pad_job_per_thread_hvx_circular(unsigned int nth, unsigned int ith,
struct htp_ops_context * octx = pctx->octx;
htp_pad_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
for (uint32_t dst_row = row_start; dst_row < row_end; dst_row++) {
uint32_t i1, i2, i3;
@@ -344,18 +345,17 @@ static void pad_job_per_thread_hvx_circular(unsigned int nth, unsigned int ith,
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
FARF(HIGH, "pad-hvx-circ %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
FARF(HIGH, "pad-hvx-circ %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
ith, nth,
src->ne[0], src->ne[1], src->ne[2], src->ne[3],
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
row_start, row_end,
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
row_start, row_end);
}
// ---------------------------------------------------------------------------
// HVX + DMA circular PAD kernel aligned, double-buffered
// HVX + DMA circular PAD kernel - aligned, double-buffered
// ---------------------------------------------------------------------------
static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int ith, void * data) {
@@ -364,9 +364,6 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
htp_pad_preamble;
htp_pad_dma_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
// -----------------------------------------------------------------------
// Priming phase: push 2 pairs of (dummy_dst_DMA, src_DMA) to seed the
// double-buffer pipeline. Every row is a real src DMA (no null DMAs).
@@ -390,6 +387,8 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
// Main loop: pop completed DMAs, assemble circular row in VTCM with
// aligned HVX ops, push dst DMA and prefetch src for the next+1 row.
// -----------------------------------------------------------------------
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t ir = row_start; ir < row_end; ir++) {
uint8_t * dst_spad_cur = (uint8_t *) dma_queue_pop(dma).src;
uint8_t * src_spad_cur = (uint8_t *) dma_queue_pop(dma).dst;
@@ -398,7 +397,7 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
pad_decompose_row(ir, ne1, ne2, &i1, &i2, &i3);
uint8_t * dst_ptr = (uint8_t *) dst->data + i1 * nb1 + i2 * nb2 + i3 * nb3;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
if (lp0 > 0) {
uint8_t * dst_left = dst_spad_cur;
const uint8_t * src_left = src_spad_cur + (size_t)(ne00 - (uint32_t)lp0) * type_size;
@@ -430,6 +429,7 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
}
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir);
dma_queue_push_vtcm_to_ddr(dma,
dma_make_ptr(dst_ptr, dst_spad_cur),
@@ -448,14 +448,11 @@ static void pad_job_per_thread_hvx_circular_dma(unsigned int nth, unsigned int i
dma_queue_flush(dma);
t2 = HAP_perf_get_qtimer_count();
FARF(HIGH, "pad-hvx-circ-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
FARF(HIGH, "pad-hvx-circ-dma %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
ith, nth,
src->ne[0], src->ne[1], src->ne[2], src->ne[3],
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
row_start, row_end,
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
row_start, row_end);
}
int op_pad(struct htp_ops_context * octx) {
@@ -489,19 +486,33 @@ int op_pad(struct htp_ops_context * octx) {
const uint32_t ne00 = src0->ne[0];
const uint32_t total_dst_rows = dst->ne[1] * dst->ne[2] * dst->ne[3];
const uint32_t n_threads = MIN(octx->n_threads, total_dst_rows > 0 ? total_dst_rows : 1);
const size_t dst_row_size = (size_t)ne0 * type_size;
uint32_t row_start = 0;
uint32_t nrows = total_dst_rows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, type_size, (uint32_t) dst_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_dst_rows, rows_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 = octx->n_threads;
const size_t src_row_size = (size_t)ne00 * type_size;
const size_t dst_row_size = (size_t)ne0 * type_size;
const size_t src_row_size_aligned = hex_round_up(src_row_size, VLEN);
const size_t dst_row_size_aligned = hex_round_up(dst_row_size, VLEN);
// Total VTCM needed: 2 buffers (ping+pong) for src and dst, per thread
const size_t vtcm_needed = (size_t)n_threads * 2 * (src_row_size_aligned + dst_row_size_aligned);
const int use_dma = (src0->nb[0] == (uint32_t)type_size) &&
(ne00 >= 512) &&
(octx->ctx->vtcm_base != NULL) &&
const int use_dma = (src0->nb[0] == (uint32_t)type_size) && (ne00 >= 512) &&
(octx->ctx->vtcm_size >= vtcm_needed);
if (use_dma) {
@@ -521,8 +532,9 @@ int op_pad(struct htp_ops_context * octx) {
.lp1 = lp1, .rp1 = rp1,
.lp2 = lp2, .rp2 = rp2,
.lp3 = lp3, .rp3 = rp3,
.nrows_per_thread = (total_dst_rows + n_threads - 1) / n_threads,
.total_dst_rows = total_dst_rows,
.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
.total_dst_rows = nrows,
.row_start = row_start,
.type_size = type_size,
.src_row_size = src_row_size,
.src_row_size_aligned = src_row_size_aligned,
@@ -537,11 +549,10 @@ int op_pad(struct htp_ops_context * octx) {
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
lp0, rp0, lp1, rp1, lp2, rp2, lp3, rp3);
if (circular && use_dma) { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx_circular_dma, &pctx, n_threads); }
else if (circular) { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx_circular, &pctx, n_threads); }
else if (use_dma) { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx_dma, &pctx, n_threads); }
else { worker_pool_run_func(octx->ctx->worker_pool, pad_job_per_thread_hvx, &pctx, n_threads); }
if (circular && use_dma) { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx_circular_dma, &pctx, n_threads); }
else if (circular) { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx_circular, &pctx, n_threads); }
else if (use_dma) { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx_dma, &pctx, n_threads); }
else { work_queue_run(octx->ctx->work_queue, pad_job_per_thread_hvx, &pctx, n_threads); }
return HTP_STATUS_OK;
}
+35 -14
View File
@@ -12,8 +12,10 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "htp-ctx.h"
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ops.h"
#include "htp-ops.h"
#include "htp-tensor.h"
struct htp_repeat_context {
struct htp_ops_context * octx;
@@ -25,6 +27,7 @@ struct htp_repeat_context {
uint32_t nrows_per_thread;
uint32_t total_dst_rows; // ne1 * ne2 * ne3
uint32_t row_start;
size_t type_size;
};
@@ -62,11 +65,11 @@ static void repeat_job_per_thread(unsigned int nth, unsigned int ith, void * dat
const size_t row_bytes = ne00 * rctx->type_size;
const uint32_t row_start = rctx->nrows_per_thread * ith;
const uint32_t row_end = MIN(row_start + rctx->nrows_per_thread, rctx->total_dst_rows);
const uint32_t row_start = rctx->row_start + rctx->nrows_per_thread * ith;
const uint32_t row_end = MIN(row_start + rctx->nrows_per_thread, rctx->row_start + rctx->total_dst_rows);
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
for (uint32_t dst_row = row_start; dst_row < row_end; dst_row++) {
// Decompose flat dst row index into (i1, i2, i3)
@@ -89,12 +92,12 @@ static void repeat_job_per_thread(unsigned int nth, unsigned int ith, void * dat
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
FARF(HIGH, "repeat %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u usec %u\n",
FARF(HIGH, "repeat %d/%d: (%ux%ux%ux%u) -> (%ux%ux%ux%u) rows %u:%u\n",
ith, nth, src->ne[0], src->ne[1], src->ne[2], src->ne[3],
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
row_start, row_end, (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
row_start, row_end);
}
int op_repeat(struct htp_ops_context * octx) {
@@ -119,21 +122,39 @@ int op_repeat(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
const uint32_t total_dst_rows = dst->ne[1] * dst->ne[2] * dst->ne[3];
const uint32_t n_threads = MIN(octx->n_threads, total_dst_rows);
if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
return HTP_STATUS_OK;
}
const uint32_t total_dst_rows = dst->ne[1] * dst->ne[2] * dst->ne[3];
const size_t dst_row_size = dst->ne[0] * type_size;
uint32_t row_start = 0;
uint32_t nrows = total_dst_rows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, type_size, (uint32_t) dst_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_dst_rows, rows_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 = octx->n_threads;
struct htp_repeat_context rctx = {
.octx = octx,
.nr0 = dst->ne[0] / src0->ne[0],
.nr1 = dst->ne[1] / src0->ne[1],
.nr2 = dst->ne[2] / src0->ne[2],
.nr3 = dst->ne[3] / src0->ne[3],
.nrows_per_thread = (total_dst_rows + n_threads - 1) / n_threads,
.total_dst_rows = total_dst_rows,
.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
.total_dst_rows = nrows,
.row_start = row_start,
.type_size = type_size,
};
@@ -142,7 +163,7 @@ int op_repeat(struct htp_ops_context * octx) {
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3],
rctx.nr0, rctx.nr1, rctx.nr2, rctx.nr3);
worker_pool_run_func(octx->ctx->worker_pool, repeat_job_per_thread, &rctx, n_threads);
work_queue_run(octx->ctx->work_queue, repeat_job_per_thread, &rctx, n_threads);
return HTP_STATUS_OK;
}
+34 -7
View File
@@ -80,6 +80,8 @@ struct htp_rope_context {
size_t dst_row_stride;
size_t src0_row_size_aligned;
uint32_t src0_nrows;
uint32_t row_start;
uint32_t nrows;
struct fastdiv_values div_ne2_ne1;
struct fastdiv_values div_ne1;
@@ -539,11 +541,11 @@ static void rope_job_f32(unsigned int nth, unsigned int ith, void * data) {
htp_rope_preamble;
const uint32_t src0_nrows = rctx->src0_nrows;
const uint32_t src0_nrows = rctx->nrows;
const uint32_t src0_nrows_per_thread = rctx->src0_nrows_per_thread;
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = rctx->row_start + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, rctx->row_start + src0_nrows);
// no work for this thread
if (src0_start_row >= src0_end_row) {
@@ -706,9 +708,32 @@ static int execute_op_rope_f32(struct htp_ops_context * octx) {
}
const struct htp_rope_kernel_params * kparams = (const struct htp_rope_kernel_params *) octx->kernel_params;
assert(kparams->n_threads > 0);
if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
assert(octx->ctx->vtcm_size >= kparams->vtcm_size);
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const size_t dst_data_row_size = dst->ne[0] * sizeof(float);
uint32_t row_start = 0;
uint32_t nrows = total_rows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_data_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
total_rows, rows_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 = octx->n_threads;
const uint32_t ne0 = dst->ne[0];
const size_t src0_row_size = src0->ne[0] * sizeof(float);
const size_t src0_row_stride = src0->nb[1];
@@ -752,15 +777,17 @@ static int execute_op_rope_f32(struct htp_ops_context * octx) {
rctx.dst_row_stride = dst_row_stride;
rctx.src0_row_size_aligned = kparams->src0_row_size_aligned;
rctx.src0_nrows = kparams->src0_nrows;
rctx.src0_nrows_per_thread = kparams->src0_nrows_per_thread;
rctx.src0_nrows = nrows;
rctx.nrows = nrows;
rctx.row_start = row_start;
rctx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
rctx.div_ne2_ne1 = kparams->div_ne2_ne1;
rctx.div_ne1 = kparams->div_ne1;
FARF(HIGH, "rope-f32 n-rows %u n-dims %d ne0 %u ext-factor %.6f theta-scale %.6f attn-factor %.6f\n", rctx.src0_nrows, rctx.n_dims, ne0,
rctx.ext_factor, rctx.theta_scale, rctx.attn_factor);
work_queue_run(octx->ctx->work_queue, rope_job_f32, &rctx, kparams->n_threads);
work_queue_run(octx->ctx->work_queue, rope_job_f32, &rctx, n_threads);
return err;
}
+42 -12
View File
@@ -18,6 +18,7 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
@@ -58,6 +59,9 @@ struct set_rows_context {
const struct htp_set_rows_kernel_params * kparams;
struct htp_set_rows_vtcm_layout vtcm_layout;
uint8_t * vtcm_base;
uint32_t task_start;
uint32_t tasks;
uint32_t tasks_per_thread;
};
#define SET_ROWS_THREAD_DMA_FN(TYPE_NAME, IDX_TYPE, COMPUTE_EXPR) \
@@ -67,12 +71,12 @@ static void set_rows_thread_dma_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsig
const struct htp_set_rows_kernel_params * kparams = srctx->kparams; \
set_rows_preamble; \
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
const uint32_t dr = kparams->tasks_per_thread; \
const uint32_t ir0 = dr * ith; \
if (ir0 >= kparams->total_tasks) { \
const uint32_t dr = srctx->tasks_per_thread; \
const uint32_t ir0 = srctx->task_start + dr * ith; \
if (ir0 >= srctx->task_start + srctx->tasks) { \
return; \
} \
const uint32_t ir1 = MIN(ir0 + dr, kparams->total_tasks); \
const uint32_t ir1 = MIN(ir0 + dr, srctx->task_start + srctx->tasks); \
dma_queue * dma_queue = octx->ctx->dma[ith]; \
const struct htp_set_rows_vtcm_layout * vtcm_layout = &srctx->vtcm_layout; \
uint8_t * vtcm_src0 = srctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread; \
@@ -192,18 +196,44 @@ int op_set_rows(struct htp_ops_context * octx) {
return HTP_STATUS_NO_SUPPORT;
}
if (octx->src[1]->type != HTP_TYPE_I32 && octx->src[1]->type != HTP_TYPE_I64) {
return HTP_STATUS_NO_SUPPORT;
if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
return HTP_STATUS_OK;
}
const struct htp_tensor * dst = octx->dst;
const uint32_t total_tasks = kparams->total_tasks;
uint32_t task_start = 0;
uint32_t tasks = total_tasks;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) && (dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0 && !htp_tensor_is_permuted(dst);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_tasks, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
task_start = range.start;
tasks = range.count;
}
if (tasks == 0) {
return HTP_STATUS_OK;
}
if (!htp_ops_context_set_n_threads(octx, (uint32_t) kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
const uint32_t n_threads = octx->n_threads;
// l2fetch the src1 (indices) tensor in the main thread
hex_l2fetch_block((const void *)octx->src[1]->data, octx->src[1]->ne[3] * octx->src[1]->nb[3]);
struct set_rows_context srctx;
srctx.octx = octx;
srctx.kparams = kparams;
srctx.task_start = task_start;
srctx.tasks = tasks;
srctx.tasks_per_thread = fastdiv(tasks + n_threads - 1, &octx->n_threads_div);
htp_set_rows_vtcm_layout_build(&srctx.vtcm_layout, octx->dst->type, ne00, kparams->n_threads);
htp_set_rows_vtcm_layout_build(&srctx.vtcm_layout, octx->dst->type, ne00, n_threads);
srctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base;
work_queue_func_t q_func = NULL;
@@ -216,15 +246,15 @@ int op_set_rows(struct htp_ops_context * octx) {
default: return HTP_STATUS_NO_SUPPORT;
}
FARF(HIGH, "set-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu n_threads %d\n",
FARF(HIGH, "set-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu 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],
srctx.vtcm_layout.src0_bytes_per_thread * kparams->n_threads,
srctx.vtcm_layout.dst_bytes_per_thread * kparams->n_threads,
kparams->n_threads);
srctx.vtcm_layout.src0_bytes_per_thread * n_threads,
srctx.vtcm_layout.dst_bytes_per_thread * n_threads,
n_threads);
work_queue_run(octx->ctx->work_queue, q_func, &srctx, kparams->n_threads);
work_queue_run(octx->ctx->work_queue, q_func, &srctx, n_threads);
return HTP_STATUS_OK;
}
+38 -15
View File
@@ -14,9 +14,11 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#define htp_softmax_preamble3 \
const uint32_t ne00 = src0->ne[0]; \
@@ -69,6 +71,8 @@ struct htp_softmax_context {
struct fastdiv_values fastdiv_ne13; // For mask broadcasting
uint32_t src0_nrows_per_thread;
uint32_t row_start;
uint32_t nrows;
};
static void apply_mask(float * restrict wp0,
@@ -223,19 +227,17 @@ static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {
htp_softmax_preamble3;
const uint32_t src0_nrows = ne01 * ne02 * ne03; // src0 rows
const uint32_t src0_nrows = smctx->nrows;
const uint32_t src0_nrows_per_thread = smctx->src0_nrows_per_thread;
const uint32_t src0_start_row = src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows);
const uint32_t src0_start_row = smctx->row_start + src0_nrows_per_thread * ith;
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, smctx->row_start + src0_nrows);
// no work for this thread
if (src0_start_row >= src0_end_row) {
return;
}
uint64_t qt = HAP_perf_get_qtimer_count();
int is_aligned = 1;
int opt_path = 0;
@@ -262,6 +264,9 @@ static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {
uint32_t prev_i2 = (uint32_t)-1;
float slope = 1.0f;
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, src0_start_row);
for (uint32_t r = src0_start_row; r < src0_end_row; ++r) {
uint32_t i1 = fastmodulo(r, ne01, &smctx->fastdiv_ne01);
uint32_t r_div_ne01 = fastdiv(r, &smctx->fastdiv_ne01);
@@ -323,10 +328,11 @@ static void softmax_job_f32(unsigned int nth, unsigned int ith, void * data) {
}
}
qt = HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - qt);
FARF(HIGH, "softmax-f32 %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u : opt %u f16 %u usec %u\n", ith, nth,
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, src0_start_row);
FARF(HIGH, "softmax-f32 %d/%d: %ux%ux%ux%u (%u:%u) x %ux%ux%ux%u -> %ux%ux%ux%u : opt %u f16 %u\n", ith, nth,
ne00, ne01, ne02, ne03, src0_start_row, src0_end_row, ne10, ne11, ne12, ne13,
ne0, ne1, ne2, ne3, opt_path, smctx->use_f16, (unsigned) qt);
ne0, ne1, ne2, ne3, opt_path, smctx->use_f16);
}
static int execute_op_softmax_f32(struct htp_ops_context * octx) {
@@ -342,13 +348,32 @@ static int execute_op_softmax_f32(struct htp_ops_context * octx) {
init_softmax_ctx(&smctx, octx);
const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads = MIN(octx->n_threads, src0_nrows);
const size_t elem_size = sizeof(float);
const size_t dst_row_size = dst->nb[1];
smctx.src0_nrows_per_thread = (src0_nrows + n_threads - 1) / n_threads;
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_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 = octx->n_threads;
smctx.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
smctx.row_start = row_start;
smctx.nrows = nrows;
const size_t src0_row_size = src0->nb[1];
const size_t src1_row_size = src0_row_size;
const size_t dst_row_size = dst->nb[1];
// VTCM scratchpads for all tensors
// 4 rows per thread, padded to HVX vector size
@@ -383,9 +408,7 @@ static int execute_op_softmax_f32(struct htp_ops_context * octx) {
octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size; octx->src1_spad.src = NULL;
octx->dst_spad.data = octx->src1_spad.data + octx->src1_spad.size; octx->dst_spad.src = NULL;
if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) return err;
worker_pool_run_func(octx->ctx->worker_pool, softmax_job_f32, &smctx, n_threads);
work_queue_run(octx->ctx->work_queue, softmax_job_f32, &smctx, n_threads);
return err;
}
+62 -30
View File
@@ -1,13 +1,16 @@
#pragma clang diagnostic ignored "-Wunused-but-set-variable"
#include <HAP_farf.h>
#include <HAP_perf.h>
#include <string.h>
#include "hex-common.h"
#include "hex-profile.h"
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#include "hvx-types.h"
#include "hvx-utils.h"
@@ -15,6 +18,7 @@ struct htp_solve_tri_context {
struct htp_ops_context * octx;
uint32_t jobs_per_thread;
uint32_t total_jobs;
uint32_t job_start;
uint32_t k_chunks;
uint32_t col_block;
};
@@ -89,11 +93,11 @@ static void solve_tri_batch_thread_f32(unsigned int nth, unsigned int ith, void
const uint32_t col_block = VLEN_FP32;
const uint32_t k_full = (k / col_block) * col_block;
const uint32_t start_batch = sctx->jobs_per_thread * ith;
const uint32_t end_batch = MIN(start_batch + sctx->jobs_per_thread, sctx->total_jobs);
const uint32_t start_batch = sctx->job_start + sctx->jobs_per_thread * ith;
const uint32_t end_batch = MIN(start_batch + sctx->jobs_per_thread, sctx->job_start + sctx->total_jobs);
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_batch);
for (uint32_t batch = start_batch; batch < end_batch; ++batch) {
const uint32_t i03 = batch / ne02;
@@ -127,11 +131,10 @@ static void solve_tri_batch_thread_f32(unsigned int nth, unsigned int ith, void
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) end_batch);
FARF(HIGH, "solve-tri-batch %d/%d: A=(%ux%u) B=(%ux%u) batch %u:%u usec %u\n",
ith, nth, n, n, k, n, start_batch, end_batch,
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
FARF(HIGH, "solve-tri-batch %d/%d: A=(%ux%u) B=(%ux%u) batch %u:%u\n",
ith, nth, n, n, k, n, start_batch, end_batch);
}
// Chunk-level thread: each job is one (batch, col_chunk) pair.
@@ -148,11 +151,11 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void
const uint32_t ne02 = src0->ne[2];
const uint32_t start_job = sctx->jobs_per_thread * ith;
const uint32_t end_job = MIN(start_job + sctx->jobs_per_thread, sctx->total_jobs);
const uint32_t start_job = sctx->job_start + sctx->jobs_per_thread * ith;
const uint32_t end_job = MIN(start_job + sctx->jobs_per_thread, sctx->job_start + sctx->total_jobs);
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_job);
for (uint32_t job = start_job; job < end_job; ++job) {
const uint32_t batch = job / sctx->k_chunks;
@@ -161,16 +164,14 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void
const uint32_t i03 = batch / ne02;
const uint32_t i02 = batch - i03 * ne02;
const uint32_t col0 = chunk * sctx->col_block;
const uint32_t coln = MIN(sctx->col_block, k - col0);
const float * A_batch =
(const float *) ((const uint8_t *) (uintptr_t) src0->data + i02 * src0->nb[2] + i03 * src0->nb[3]);
const float * B_batch =
(const float *) ((const uint8_t *) (uintptr_t) src1->data + i02 * src1->nb[2] + i03 * src1->nb[3]);
float * X_batch = (float *) ((uint8_t *) (uintptr_t) dst->data + i02 * dst->nb[2] + i03 * dst->nb[3]);
const bool use_hvx = (coln >= 8);
const uint32_t col0 = chunk * sctx->col_block;
const uint32_t coln = MIN(sctx->col_block, k - col0);
for (uint32_t row = 0; row < n; ++row) {
const float diag = A_batch[row * n + row];
@@ -179,7 +180,7 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void
const float * A_row = A_batch + row * n;
const float * B_row = B_batch + row * k;
if (use_hvx) {
if (coln >= 8) {
solve_tri_row_hvx(A_row, B_row, X_batch, row, k, col0, coln, inv_diag);
} else {
solve_tri_row_scalar(A_row, B_row, X_batch, row, k, col0, coln, inv_diag);
@@ -187,11 +188,10 @@ static void solve_tri_chunk_thread_f32(unsigned int nth, unsigned int ith, void
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) end_job);
FARF(HIGH, "solve-tri-chunk %d/%d: A=(%ux%u) B=(%ux%u) job %u:%u usec %u\n",
ith, nth, n, n, k, n, start_job, end_job,
(unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
FARF(HIGH, "solve-tri-chunk %d/%d: A=(%ux%u) B=(%ux%u) jobs %u:%u\n",
ith, nth, n, n, k, n, start_job, end_job);
}
int op_solve_tri(struct htp_ops_context * octx) {
@@ -235,32 +235,64 @@ int op_solve_tri(struct htp_ops_context * octx) {
dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3], batched);
if (batched) {
uint32_t job_start = 0;
uint32_t njobs = total_batches;
if (octx->ctx->mdev.count > 1) {
const uint32_t batch_size = dst->nb[2];
const uint32_t batches_per_chunk = (batch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(batch_size, HEX_L2_LINE_SIZE)) : 1;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_batches, htp_tensor_mdev_data_aligned(dst) ? batches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
job_start = range.start;
njobs = range.count;
}
if (njobs == 0) {
return HTP_STATUS_OK;
}
// Batch-level parallelism
const uint32_t n_threads = MIN((uint32_t) octx->n_threads, total_batches);
const uint32_t n_threads = octx->n_threads;
struct htp_solve_tri_context sctx = {
.octx = octx,
.jobs_per_thread = (total_batches + n_threads - 1) / n_threads,
.total_jobs = total_batches,
.jobs_per_thread = fastdiv(njobs + n_threads - 1, &octx->n_threads_div),
.total_jobs = njobs,
.job_start = job_start,
.k_chunks = k_chunks,
.col_block = col_block,
};
worker_pool_run_func(octx->ctx->worker_pool, solve_tri_batch_thread_f32, &sctx, n_threads);
work_queue_run(octx->ctx->work_queue, solve_tri_batch_thread_f32, &sctx, n_threads);
} else {
// Chunk-level parallelism
const uint32_t total_jobs = total_batches * k_chunks;
const uint32_t n_threads = MIN((uint32_t) octx->n_threads, MAX(total_jobs, 1));
uint32_t job_start = 0;
uint32_t njobs = total_jobs;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) && ((dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_jobs, can_split ? 1 : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
job_start = range.start;
njobs = range.count;
}
if (njobs == 0) {
return HTP_STATUS_OK;
}
const uint32_t n_threads = octx->n_threads;
struct htp_solve_tri_context sctx = {
.octx = octx,
.jobs_per_thread = (total_jobs + n_threads - 1) / n_threads,
.total_jobs = total_jobs,
.jobs_per_thread = fastdiv(njobs + n_threads - 1, &octx->n_threads_div),
.total_jobs = njobs,
.job_start = job_start,
.k_chunks = k_chunks,
.col_block = col_block,
};
worker_pool_run_func(octx->ctx->worker_pool, solve_tri_chunk_thread_f32, &sctx, n_threads);
work_queue_run(octx->ctx->work_queue, solve_tri_chunk_thread_f32, &sctx, n_threads);
}
return HTP_STATUS_OK;
+96 -73
View File
@@ -4,7 +4,6 @@
#include <HAP_farf.h>
#include <HAP_mem.h>
#include <HAP_perf.h>
#include <HAP_ps.h>
#include <hexagon_protos.h>
#include <hexagon_types.h>
@@ -16,8 +15,9 @@
#include "ggml-common.h"
#include "htp-ctx.h"
#include "hex-dma.h"
#include "hex-profile.h"
#include "htp-ops.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#include "hvx-utils.h"
#define htp_ssm_conv_tensors_preamble \
@@ -63,6 +63,8 @@ struct htp_ssm_conv_context {
uint32_t nrows_per_thread;
uint32_t d_inner_tile;
uint64_t t_start;
uint32_t row_start;
uint32_t nrows;
};
#define htp_ssm_conv_preamble \
@@ -75,9 +77,6 @@ struct htp_ssm_conv_context {
static void ssm_conv_thread_f32_f32(unsigned int nth, unsigned int ith, void *data) {
htp_ssm_conv_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
const uint32_t d_conv = src1->ne[0];
const uint32_t d_inner = src0->ne[1];
const uint32_t n_t = dst->ne[1];
@@ -95,14 +94,17 @@ static void ssm_conv_thread_f32_f32(unsigned int nth, unsigned int ith, void *da
// Calculate row range for this thread
const uint32_t d_inner_per_thread = scctx->nrows_per_thread;
const uint32_t d_inner_start = d_inner_per_thread * ith;
const uint32_t d_inner_end = MIN(d_inner_start + d_inner_per_thread, d_inner);
const uint32_t d_inner_start = scctx->row_start + d_inner_per_thread * ith;
const uint32_t d_inner_end = MIN(d_inner_start + d_inner_per_thread, scctx->row_start + scctx->nrows);
// No work for this thread
if (d_inner_start >= d_inner_end) {
return;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) d_inner_start);
for (uint32_t i3 = 0; i3 < n_s; ++i3) {
for (uint32_t i2 = 0; i2 < n_t; ++i2) {
for (uint32_t i1 = d_inner_start; i1 < d_inner_end; ++i1) {
@@ -121,12 +123,12 @@ static void ssm_conv_thread_f32_f32(unsigned int nth, unsigned int ith, void *da
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) d_inner_end);
FARF(HIGH, "ssm-conv-f32 %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n",
FARF(HIGH, "ssm-conv-f32 %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u\n",
ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], d_inner_start, d_inner_end,
src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0], dst->ne[1],
dst->ne[2], dst->ne[3], (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
dst->ne[2], dst->ne[3]);
}
@@ -257,9 +259,6 @@ static inline void transpose_src0_block(const float * src0_block,
static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void *data) {
htp_ssm_conv_preamble;
uint64_t t1, t2;
t1 = HAP_perf_get_qtimer_count();
const uint32_t d_conv = src1->ne[0];
const uint32_t d_inner = src0->ne[1];
const uint32_t n_t = dst->ne[1];
@@ -273,13 +272,16 @@ static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void
const uint32_t dst_stride_seq = dst->nb[2] / sizeof(float);
const uint32_t dr = scctx->nrows_per_thread;
const uint32_t ir0 = dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, d_inner);
const uint32_t ir0 = scctx->row_start + dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, scctx->row_start + scctx->nrows);
if (ir0 >= ir1) {
return;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir0);
const uint32_t d_inner_per_thread = ir1 - ir0;
const uint32_t d_inner_stride = scctx->nrows_per_thread;
const uint32_t d_inner_tile = scctx->d_inner_tile;
@@ -319,97 +321,118 @@ static void ssm_conv_thread_f32_f32_hvx(unsigned int nth, unsigned int ith, void
HVX_Vector w = *(const HVX_Vector *) (src1_T + j * d_inner_stride + tile_off + cb);
acc = Q6_Vqf32_vadd_Vqf32Vqf32(acc, Q6_Vqf32_vmpy_VsfVsf(x, w));
}
HVX_Vector res = Q6_Vsf_equals_Vqf32(acc);
float * dst_ptr = dst_data + i3 * dst_stride_seq + t * dst_stride_token + (ir0 + tile_off + cb);
HVX_Vector y = Q6_Vsf_equals_Vqf32(acc);
float * dst_ptr = dst_data + (ir0 + tile_off + cb) + t * dst_stride_token + i3 * dst_stride_seq;
if (cb_n == C_TILE) {
*(HVX_UVector *) dst_ptr = res;
*(HVX_UVector *) dst_ptr = y;
} else {
hvx_vec_store_u(dst_ptr, cb_n * sizeof(float), res);
hvx_vec_store_u(dst_ptr, cb_n * sizeof(float), y);
}
}
}
}
}
t2 = HAP_perf_get_qtimer_count();
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) ir1);
FARF(HIGH, "ssm-conv-f32-hvx %d/%d: %ux%ux%ux%u (%u:%u) tile=%u * %ux%ux%ux%u -> %ux%ux%ux%u usec %u\n",
ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1, d_inner_tile,
FARF(HIGH, "ssm-conv-f32-hvx %d/%d: %ux%ux%ux%u (%u:%u) * %ux%ux%ux%u -> %ux%ux%ux%u\n",
ith, nth, src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], ir0, ir1,
src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0], dst->ne[1],
dst->ne[2], dst->ne[3], (unsigned) HAP_perf_qtimer_count_to_us(t2 - t1));
dst->ne[2], dst->ne[3]);
}
int op_ssm_conv_f32(struct htp_ops_context * octx) {
htp_ssm_conv_tensors_preamble;
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * src1 = octx->src[1];
const struct htp_tensor * dst = octx->dst;
if (src0->type != HTP_TYPE_F32 || src1->type != HTP_TYPE_F32 || dst->type != HTP_TYPE_F32) {
FARF(ERROR, "ssm_conv: only (F32 x F32 -> F32) OPs supported");
return HTP_STATUS_NO_SUPPORT;
}
struct htp_ssm_conv_context scctx = { 0 };
scctx.octx = octx;
const uint32_t d_conv = src1->ne[0];
const uint32_t d_inner = src0->ne[1];
const uint32_t n_t = dst->ne[1]; // tokens per sequence
const uint32_t n_s = dst->ne[2]; // number of sequences in the batch
const uint32_t n_threads = MIN(octx->n_threads, d_inner);
if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
return HTP_STATUS_OK;
}
if (!(octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
uint32_t use_hvx = 0;
if (d_inner >= VLEN_FP32 && n_t >= VLEN_FP32) {
use_hvx = 1;
}
uint32_t row_start = 0;
uint32_t nrows = d_inner;
scctx.nrows_per_thread = hex_round_up((d_inner + n_threads - 1) / n_threads, VLEN_FP32);
if (octx->ctx->mdev.count > 1) {
const uint32_t elems_per_chunk = VLEN_FP32;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(d_inner, htp_tensor_mdev_data_aligned(dst) ? elems_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
row_start = range.start;
nrows = range.count;
}
const uint32_t d_inner_per_thread = scctx.nrows_per_thread;
const uint32_t ncs = src0->ne[0];
if (nrows == 0) {
return HTP_STATUS_OK;
}
const uint32_t src1_T_size = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 256);
const uint32_t src0_T_max = HTP_SSM_CONV_VTCM_BUDGET > src1_T_size ? HTP_SSM_CONV_VTCM_BUDGET - src1_T_size : 0;
const uint32_t n_threads = octx->n_threads;
uint32_t d_inner_tile = (src0_T_max / sizeof(float)) / ncs;
d_inner_tile -= (d_inner_tile % VLEN_FP32);
if (d_inner_tile == 0) {
FARF(HIGH, "ssm_conv-f32: inner tile rounds to 0 (ncs=%u), falling back to scalar\n", ncs);
struct htp_ssm_conv_context scctx = { 0 };
scctx.octx = octx;
scctx.row_start = row_start;
scctx.nrows = nrows;
uint32_t use_hvx = 0;
if (nrows >= VLEN_FP32 && n_t >= VLEN_FP32) {
use_hvx = 1;
}
const uint32_t raw_rpt = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
scctx.nrows_per_thread = hex_round_up(raw_rpt, VLEN_FP32);
const uint32_t d_inner_per_thread = scctx.nrows_per_thread;
const uint32_t ncs = src0->ne[0];
const uint32_t src1_T_size = hex_round_up(d_conv * d_inner_per_thread * sizeof(float), 256);
const uint32_t src0_T_max = HTP_SSM_CONV_VTCM_BUDGET > src1_T_size ? HTP_SSM_CONV_VTCM_BUDGET - src1_T_size : 0;
uint32_t d_inner_tile = (src0_T_max / sizeof(float)) / ncs;
d_inner_tile -= (d_inner_tile % VLEN_FP32);
if (d_inner_tile == 0) {
FARF(HIGH, "ssm_conv-f32: inner tile rounds to 0 (ncs=%u), falling back to scalar\n", ncs);
use_hvx = 0;
} else {
scctx.d_inner_tile = d_inner_tile;
octx->src0_spad.size_per_thread = hex_round_up(d_inner_tile * ncs * sizeof(float), 256);
octx->src1_spad.size_per_thread = src1_T_size;
octx->dst_spad.size_per_thread = 0;
octx->src0_spad.size = octx->src0_spad.size_per_thread * n_threads;
octx->src1_spad.size = octx->src1_spad.size_per_thread * n_threads;
octx->dst_spad.size = 0;
octx->src0_spad.data = octx->ctx->vtcm_base;
octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size;
octx->src0_spad.src = NULL;
octx->src1_spad.src = NULL;
const size_t total_spad = octx->src0_spad.size + octx->src1_spad.size;
if (total_spad > octx->ctx->vtcm_size) {
FARF(HIGH, "ssm_conv-f32: scratchpad %zu exceeds VTCM %zu, falling back to scalar\n",
total_spad, octx->ctx->vtcm_size);
use_hvx = 0;
} else {
scctx.d_inner_tile = d_inner_tile;
octx->src0_spad.size_per_thread = hex_round_up(d_inner_tile * ncs * sizeof(float), 256);
octx->src1_spad.size_per_thread = src1_T_size;
octx->dst_spad.size_per_thread = 0;
octx->src0_spad.size = octx->src0_spad.size_per_thread * n_threads;
octx->src1_spad.size = octx->src1_spad.size_per_thread * n_threads;
octx->dst_spad.size = 0;
octx->src0_spad.data = octx->ctx->vtcm_base;
octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size;
octx->src0_spad.src = NULL;
octx->src1_spad.src = NULL;
const size_t total_spad = octx->src0_spad.size + octx->src1_spad.size;
if (total_spad > octx->ctx->vtcm_size) {
FARF(HIGH, "ssm_conv-f32: scratchpad %zu exceeds VTCM %zu, falling back to scalar\n",
total_spad, octx->ctx->vtcm_size);
use_hvx = 0;
}
}
}
FARF(HIGH, "ssm-conv-f32: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_hvx %d\n", src0->ne[0],
src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0],
dst->ne[1], dst->ne[2], dst->ne[3], use_hvx);
FARF(HIGH, "ssm-conv-f32: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : use_hvx %d\n", src0->ne[0],
src0->ne[1], src0->ne[2], src0->ne[3], src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3], dst->ne[0],
dst->ne[1], dst->ne[2], dst->ne[3], use_hvx);
if (use_hvx) {
worker_pool_run_func(octx->ctx->worker_pool, ssm_conv_thread_f32_f32_hvx, &scctx, n_threads);
} else {
worker_pool_run_func(octx->ctx->worker_pool, ssm_conv_thread_f32_f32, &scctx, n_threads);
}
if (use_hvx) {
work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_f32_hvx, &scctx, n_threads);
} else {
work_queue_run(octx->ctx->work_queue, ssm_conv_thread_f32_f32, &scctx, n_threads);
}
return HTP_STATUS_OK;
+53 -27
View File
@@ -13,35 +13,38 @@
#define GGML_COMMON_DECL_C
#include "ggml-common.h"
#include "hex-common.h"
#include "hex-profile.h"
#include "htp-ctx.h"
#include "htp-ops.h"
#include "htp-ops.h"
#include "htp-tensor.h"
#define sum_rows_preamble \
const struct htp_tensor *src0 = octx->src[0]; \
const struct htp_tensor *dst = octx->dst; \
\
const uint32_t ne00 = src0->ne[0]; \
const uint32_t ne01 = src0->ne[1]; \
const uint32_t ne02 = src0->ne[2]; \
const uint32_t ne03 = src0->ne[3]; \
\
const uint32_t nb00 = src0->nb[0]; \
const uint32_t nb01 = src0->nb[1]; \
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb0 = dst->nb[0]; \
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3]; \
const uint32_t ne00 = src0->ne[0]; \
const uint32_t ne01 = src0->ne[1]; \
const uint32_t ne02 = src0->ne[2]; \
const uint32_t ne03 = src0->ne[3]; \
\
const uint32_t nb00 = src0->nb[0]; \
const uint32_t nb01 = src0->nb[1]; \
const uint32_t nb02 = src0->nb[2]; \
const uint32_t nb03 = src0->nb[3]; \
\
const uint32_t ne0 = dst->ne[0]; \
const uint32_t ne1 = dst->ne[1]; \
const uint32_t ne2 = dst->ne[2]; \
const uint32_t ne3 = dst->ne[3]; \
\
const uint32_t nb0 = dst->nb[0]; \
const uint32_t nb1 = dst->nb[1]; \
const uint32_t nb2 = dst->nb[2]; \
const uint32_t nb3 = dst->nb[3]; \
struct sum_rows_context {
struct htp_ops_context * octx;
const uint8_t * src_data;
uint8_t * dst_data;
uint32_t ne00;
@@ -76,6 +79,9 @@ static void sum_rows_thread_f32(unsigned int nth, unsigned int ith, void *data)
// Calculate actual number of rows for this thread
const uint32_t n_rows = end_row - start_row;
struct htp_thread_trace * tr = &smctx->octx->ctx->trace[ith];
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)));
@@ -89,6 +95,8 @@ static void sum_rows_thread_f32(unsigned int nth, unsigned int ith, void *data)
dst_th[ir] = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_row);
}
int op_sum_rows(struct htp_ops_context * octx) {
@@ -102,9 +110,26 @@ int op_sum_rows(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
const uint32_t src0_nrows = ne01 * ne02 * ne03;
const uint32_t n_threads = MIN(octx->n_threads, src0_nrows);
const uint32_t rows_per_thread = (src0_nrows + n_threads - 1) / n_threads;
const uint32_t src0_nrows = ne01 * ne02 * ne03;
const size_t dst_data_row_size = dst->ne[0] * sizeof(float);
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_data_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_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 = octx->n_threads;
const uint32_t rows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div);
bool opt_path = false;
if ((0 == hex_is_aligned((void *) src0->data, VLEN)) && !(nb01 & (VLEN - 1))) {
@@ -112,17 +137,18 @@ int op_sum_rows(struct htp_ops_context * octx) {
}
struct sum_rows_context smctx = {
.src_data = (const uint8_t *) src0->data,
.dst_data = (uint8_t *) dst->data,
.octx = octx,
.src_data = (const uint8_t *) src0->data + row_start * nb01,
.dst_data = (uint8_t *) dst->data + row_start * nb1,
.ne00 = ne00,
.src_stride = nb01,
.dst_stride = nb1,
.rows_per_thread = rows_per_thread,
.total_rows = src0_nrows,
.total_rows = nrows,
.opt_path = opt_path,
};
worker_pool_run_func(octx->ctx->worker_pool, sum_rows_thread_f32, &smctx, n_threads);
work_queue_run(octx->ctx->work_queue, sum_rows_thread_f32, &smctx, n_threads);
return HTP_STATUS_OK;
}
+150 -128
View File
@@ -46,6 +46,7 @@ struct htp_unary_context {
uint32_t block;
uint32_t src0_nrows;
uint32_t src0_nrows_per_thread;
uint32_t row_start;
uint32_t nc;
uint32_t col_tile; // tiled mode
bool broadcast_weight;
@@ -496,7 +497,7 @@ static void tri_f32(const float * restrict src,
}
if (boundary > ne0) boundary = ne0;
// Full HVX vectors each starts at a 128-byte aligned offset
// Full HVX vectors - each starts at a 128-byte aligned offset
for (uint32_t i = 0; i < nvec; i++) {
const uint32_t vec_start = i * VLEN_FP32;
const uint32_t vec_end = vec_start + VLEN_FP32;
@@ -563,7 +564,7 @@ static void softplus_f32(const float * restrict src,
for (uint32_t i = 0; i < ne0; i++) {
float x = src_f[i];
// For x > 20: softplus(x) x (avoids exp overflow)
// For x > 20: softplus(x) ~ x (avoids exp overflow)
dst_f[i] = (x > 20.0f) ? x : logf(1.0f + expf(x));
}
}
@@ -661,8 +662,8 @@ static void unary_task_##SUFFIX##_##NAME(unsigned int nth, unsigned int ith, voi
const size_t dst_row_size_aligned = uctx->dst_row_size_aligned; \
\
const uint32_t src0_nrows = uctx->src0_nrows; \
const uint32_t src0_start_row = src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, src0_nrows); \
const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows); \
\
if (src0_start_row >= src0_end_row) { \
return; \
@@ -833,124 +834,126 @@ DEFINE_UNARY_TASK_IMPL(unary_abs, _Float16, f16, false, false, abs_f16(src0_vtcm
DEFINE_UNARY_TASK_IMPL(unary_log, _Float16, f16, false, false, log_f16(src0_vtcm, dst_vtcm, block_size, uctx))
// Apply a pointwise unary op to one column tile that is already in VTCM.
#define DEFINE_UNARY_TILED_TASK(NAME, IS_TRI, CORE_TILE_EXPR) \
static void unary_task_f32_tiled_##NAME(unsigned int nth, unsigned int ith, void * data) { \
const struct htp_unary_context * uctx = (const struct htp_unary_context *) data; \
struct htp_ops_context * octx = uctx->octx; \
const struct htp_tensor * src = octx->src[0]; \
const struct htp_tensor * dst = octx->dst; \
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
\
htp_unary_preamble; \
\
int32_t * op_params = octx->op_params; \
const uint32_t col_tile = uctx->col_tile; \
\
const uint32_t src0_nrows = uctx->src0_nrows; \
const uint32_t src0_start_row = uctx->src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + uctx->src0_nrows_per_thread, src0_nrows); \
\
if (src0_start_row >= src0_end_row) { \
return; \
} \
\
const uint8_t * restrict data_src = uctx->data_src0; \
uint8_t * restrict data_dst = uctx->data_dst; \
\
uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread); \
uint8_t * dst_vtcm_data = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread); \
\
const size_t src0_half = uctx->src0_vtcm_half_size; \
const size_t dst_half = uctx->dst_vtcm_half_size; \
\
dma_queue * dmaq = octx->ctx->dma[ith]; \
\
const struct fastdiv_values * div_ne01 = &uctx->kparams->div_ne01; \
const struct fastdiv_values * div_ne02 = &uctx->kparams->div_ne02; \
const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012; \
const struct fastdiv_values * div_tpr = &uctx->kparams->div_tpr; \
\
const uint32_t tiles_per_row = (ne0 + col_tile - 1) / col_tile; \
const int32_t tri_ttype = (IS_TRI) ? op_params[0] : 0; \
\
const bool src0_contig = (nb02 == (size_t)ne01 * nb01) && \
(nb03 == (size_t)ne02 * nb02); \
const bool dst_contig = (nb2 == (size_t)ne1 * nb1) && \
(nb3 == (size_t)ne2 * nb2); \
\
const uint32_t total_tiles = (src0_end_row - src0_start_row) * tiles_per_row; \
\
for (uint32_t t = 0, vtcm_idx = 0; t < total_tiles && vtcm_idx < 2; t++, vtcm_idx++) { \
const uint32_t row = src0_start_row + t / tiles_per_row; \
const uint32_t col = (t % tiles_per_row) * col_tile; \
const uint32_t tw = MIN(col_tile, ne0 - col); \
const size_t tb = (size_t) tw * sizeof(float); \
const size_t soff = (src0_contig ? (row * nb01) : \
unary_row_offset(row, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03)) +\
(size_t) col * sizeof(float); \
\
dma_queue_push(dmaq, dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_half)), 0, 0, 0, 0); \
dma_queue_push(dmaq, dma_make_ptr(src0_vtcm_data + (vtcm_idx * src0_half), data_src + soff), tb, tb, tb, 1);\
} \
\
uint32_t row = src0_start_row; \
uint32_t col = 0; \
uint32_t tile_in_row = 0; \
uint32_t i01 = fastmodulo(row, ne01, div_ne01); \
\
uint32_t prow = src0_start_row + fastdiv(2, div_tpr); \
uint32_t pcol = fastmodulo(2, tiles_per_row, div_tpr) * col_tile; \
uint32_t ptile_in_row = fastmodulo(2, tiles_per_row, div_tpr); \
\
for (uint32_t t = 0; t < total_tiles; t++) { \
uint8_t * dst_vtcm = (uint8_t *) dma_queue_pop(dmaq).src; \
uint8_t * src_vtcm = (uint8_t *) dma_queue_pop(dmaq).dst; \
\
const uint32_t tw = MIN(col_tile, ne0 - col); \
\
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, t); \
CORE_TILE_EXPR; \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, t); \
\
const size_t doff = (dst_contig ? (row * nb1) : \
unary_row_offset(row, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3)) + \
(size_t) col * sizeof(float); \
const size_t tb = (size_t) tw * sizeof(float); \
dma_queue_push(dmaq, dma_make_ptr(data_dst + doff, dst_vtcm), tb, tb, tb, 1); \
\
const uint32_t pt = t + 2; \
if (pt < total_tiles) { \
const uint32_t ptw = MIN(col_tile, ne0 - pcol); \
const size_t ptb = (size_t) ptw * sizeof(float); \
const size_t psoff = (src0_contig ? (prow * nb01) : \
unary_row_offset(prow, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, \
nb03)) + \
(size_t) pcol * sizeof(float); \
dma_queue_push(dmaq, dma_make_ptr(src_vtcm, data_src + psoff), ptb, ptb, ptb, 1); \
} \
\
tile_in_row++; \
col += col_tile; \
if (tile_in_row == tiles_per_row) { \
tile_in_row = 0; \
col = 0; \
row++; \
i01++; \
if (i01 == ne01) { \
i01 = 0; \
} \
} \
\
ptile_in_row++; \
pcol += col_tile; \
if (ptile_in_row == tiles_per_row) { \
ptile_in_row = 0; \
pcol = 0; \
prow++; \
} \
} \
\
dma_queue_flush(dmaq); \
#define DEFINE_UNARY_TILED_TASK(NAME, IS_TRI, CORE_TILE_EXPR) \
static void unary_task_f32_tiled_##NAME(unsigned int nth, unsigned int ith, void * data) { \
const struct htp_unary_context * uctx = (const struct htp_unary_context *) data; \
struct htp_ops_context * octx = uctx->octx; \
const struct htp_tensor * src = octx->src[0]; \
const struct htp_tensor * dst = octx->dst; \
struct htp_thread_trace * tr = &octx->ctx->trace[ith]; \
\
htp_unary_preamble; \
\
uint32_t src0_nrows_per_thread = uctx->src0_nrows_per_thread; \
\
int32_t * op_params = octx->op_params; \
const uint32_t col_tile = uctx->col_tile; \
\
const uint32_t src0_nrows = uctx->src0_nrows; \
const uint32_t src0_start_row = uctx->row_start + src0_nrows_per_thread * ith; \
const uint32_t src0_end_row = MIN(src0_start_row + src0_nrows_per_thread, uctx->row_start + src0_nrows); \
\
if (src0_start_row >= src0_end_row) { \
return; \
} \
\
const uint8_t * restrict data_src = uctx->data_src0; \
uint8_t * restrict data_dst = uctx->data_dst; \
\
uint8_t * src0_vtcm_data = uctx->vtcm_src0 + (ith * uctx->vtcm_src0_size_per_thread); \
uint8_t * dst_vtcm_data = uctx->vtcm_dst + (ith * uctx->vtcm_dst_size_per_thread); \
\
const size_t src0_half = uctx->src0_vtcm_half_size; \
const size_t dst_half = uctx->dst_vtcm_half_size; \
\
dma_queue * dmaq = octx->ctx->dma[ith]; \
\
const struct fastdiv_values * div_ne01 = &uctx->kparams->div_ne01; \
const struct fastdiv_values * div_ne02 = &uctx->kparams->div_ne02; \
const struct fastdiv_values * div_ne012 = &uctx->kparams->div_ne012; \
const struct fastdiv_values * div_tpr = &uctx->kparams->div_tpr; \
\
const uint32_t tiles_per_row = (ne0 + col_tile - 1) / col_tile; \
const int32_t tri_ttype = (IS_TRI) ? op_params[0] : 0; \
\
const bool src0_contig = (nb02 == (size_t)ne01 * nb01) && \
(nb03 == (size_t)ne02 * nb02); \
const bool dst_contig = (nb2 == (size_t)ne1 * nb1) && \
(nb3 == (size_t)ne2 * nb2); \
\
const uint32_t total_tiles = (src0_end_row - src0_start_row) * tiles_per_row; \
\
for (uint32_t t = 0, vtcm_idx = 0; t < total_tiles && vtcm_idx < 2; t++, vtcm_idx++) { \
const uint32_t row = src0_start_row + t / tiles_per_row; \
const uint32_t col = (t % tiles_per_row) * col_tile; \
const uint32_t tw = MIN(col_tile, ne0 - col); \
const size_t tb = (size_t) tw * sizeof(float); \
const size_t soff = (src0_contig ? (row * nb01) : \
unary_row_offset(row, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, nb03)) + \
(size_t) col * sizeof(float); \
\
dma_queue_push(dmaq, dma_make_ptr(data_dst, dst_vtcm_data + (vtcm_idx * dst_half)), 0, 0, 0, 0); \
dma_queue_push(dmaq, dma_make_ptr(src0_vtcm_data + (vtcm_idx * src0_half), data_src + soff), tb, tb, tb, 1); \
} \
\
uint32_t row = src0_start_row; \
uint32_t col = 0; \
uint32_t tile_in_row = 0; \
uint32_t i01 = fastmodulo(row, ne01, div_ne01); \
\
uint32_t prow = src0_start_row + fastdiv(2, div_tpr); \
uint32_t pcol = fastmodulo(2, tiles_per_row, div_tpr) * col_tile; \
uint32_t ptile_in_row = fastmodulo(2, tiles_per_row, div_tpr); \
\
for (uint32_t t = 0; t < total_tiles; t++) { \
uint8_t * dst_vtcm = (uint8_t *) dma_queue_pop(dmaq).src; \
uint8_t * src_vtcm = (uint8_t *) dma_queue_pop(dmaq).dst; \
\
const uint32_t tw = MIN(col_tile, ne0 - col); \
\
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, t); \
CORE_TILE_EXPR; \
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, t); \
\
const size_t doff = (dst_contig ? (row * nb1) : \
unary_row_offset(row, ne1, ne2, div_ne01, div_ne02, div_ne012, nb1, nb2, nb3)) + \
(size_t) col * sizeof(float); \
const size_t tb = (size_t) tw * sizeof(float); \
dma_queue_push(dmaq, dma_make_ptr(data_dst + doff, dst_vtcm), tb, tb, tb, 1); \
\
const uint32_t pt = t + 2; \
if (pt < total_tiles) { \
const uint32_t ptw = MIN(col_tile, ne0 - pcol); \
const size_t ptb = (size_t) ptw * sizeof(float); \
const size_t psoff = (src0_contig ? (prow * nb01) : \
unary_row_offset(prow, ne01, ne02, div_ne01, div_ne02, div_ne012, nb01, nb02, \
nb03)) + \
(size_t) pcol * sizeof(float); \
dma_queue_push(dmaq, dma_make_ptr(src_vtcm, data_src + psoff), ptb, ptb, ptb, 1); \
} \
\
tile_in_row++; \
col += col_tile; \
if (tile_in_row == tiles_per_row) { \
tile_in_row = 0; \
col = 0; \
row++; \
i01++; \
if (i01 == ne01) { \
i01 = 0; \
} \
} \
\
ptile_in_row++; \
pcol += col_tile; \
if (ptile_in_row == tiles_per_row) { \
ptile_in_row = 0; \
pcol = 0; \
prow++; \
} \
} \
\
dma_queue_flush(dmaq); \
}
static inline void tile_scale_f32(uint8_t * dst_vtcm, const uint8_t * src_vtcm, uint32_t tw, const int32_t * op_params) {
@@ -1146,14 +1149,32 @@ static int execute_op_unary(struct htp_ops_context * octx) {
const struct htp_unary_kernel_params * kparams = (const struct htp_unary_kernel_params *) octx->kernel_params;
if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads = kparams->n_threads;
const size_t elem_size = is_f16 ? sizeof(_Float16) : sizeof(float);
const size_t src0_data_row_size = src0->ne[0] * elem_size;
const size_t dst_data_row_size = dst->ne[0] * elem_size;
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
if (octx->ctx->mdev.count > 1) {
uint32_t rows_per_chunk = 0;
htp_tensor_mdev_rows_per_chunk(dst, (uint32_t) elem_size, (uint32_t) dst_data_row_size, &rows_per_chunk);
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(src0_nrows, rows_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 = octx->n_threads;
const size_t src0_row_size_aligned = kparams->src0_row_size_aligned;
const size_t dst_row_size_aligned = kparams->dst_row_size_aligned;
@@ -1191,8 +1212,9 @@ static int execute_op_unary(struct htp_ops_context * octx) {
struct htp_unary_context uctx = {
.octx = octx,
.kparams = kparams,
.src0_nrows_per_thread = (src0_nrows + n_threads - 1) / n_threads,
.src0_nrows = src0_nrows,
.src0_nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
.src0_nrows = nrows,
.row_start = row_start,
.data_src0 = (const uint8_t *)src0->data,
.data_src1 = (octx->op == HTP_OP_RMS_NORM_MUL) ? (const uint8_t *)src1->data : NULL,
@@ -1287,7 +1309,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
}
if (task_func) {
worker_pool_run_func(octx->ctx->worker_pool, task_func, &uctx, n_threads);
work_queue_run(octx->ctx->work_queue, task_func, &uctx, n_threads);
} else {
FARF(ERROR, "execute_op_unary: task function is NULL for op %d\n", octx->op);
err = HTP_STATUS_NO_SUPPORT;
+1
View File
@@ -10,6 +10,7 @@ ggml_add_backend_library(ggml-metal
ggml-metal-device.cpp
ggml-metal-common.cpp
ggml-metal-context.m
ggml-metal-fusion.cpp
ggml-metal-ops.cpp
ggml-metal-tuning.cpp
)
+10 -37
View File
@@ -1,4 +1,5 @@
#include "ggml-metal-common.h"
#include "ggml-metal-fusion.h"
#include "ggml.h"
#include "ggml-impl.h"
@@ -390,59 +391,31 @@ static std::vector<int> ggml_metal_graph_optimize_reorder(const std::vector<node
}
void ggml_graph_optimize(ggml_cgraph * gf) {
constexpr int MAX_FUSE = 16;
const int n = gf->n_nodes;
enum ggml_op ops[MAX_FUSE];
std::vector<node_info> nodes;
nodes.reserve(gf->n_nodes);
// fuse nodes:
// we don't want to make reorders that break fusing, so we first pack all fusable tensors
// and perform the reorder over the fused nodes. after the reorder is done, we unfuse
//
// the fusable sequences are declared in the fusion table (ggml-metal-fuse.cpp), so the
// packing here is driven by the same patterns that the op encoders will later use
for (int i = 0; i < n; i++) {
node_info node = {
/*.node =*/ gf->nodes[i],
/*.fused =*/ {},
};
// fuse only ops that start with these operations
// can be expanded when needed
if (node.op() == GGML_OP_ADD ||
node.op() == GGML_OP_NORM ||
node.op() == GGML_OP_RMS_NORM) {
ops[0] = node.op();
const int f = ggml_metal_fusion_max(gf, i);
int f = i + 1;
while (f < n && f < i + MAX_FUSE) {
// conservatively allow fusing only these ops
// can be expanded when needed
if (gf->nodes[f]->op != GGML_OP_ADD &&
gf->nodes[f]->op != GGML_OP_MUL &&
gf->nodes[f]->op != GGML_OP_NORM &&
gf->nodes[f]->op != GGML_OP_RMS_NORM) {
break;
}
ops[f - i] = gf->nodes[f]->op;
f++;
}
// add the fused tensors into the node info so we can unfuse them later
for (int k = 1; k < f; k++) {
++i;
f -= i;
for (; f > 1; f--) {
if (ggml_can_fuse(gf, i, ops, f)) {
break;
}
}
// add the fused tensors into the node info so we can unfuse them later
for (int k = 1; k < f; k++) {
++i;
// the .dst() becomes the last fused tensor
node.add_fused(gf->nodes[i]);
}
// the .dst() becomes the last fused tensor
node.add_fused(gf->nodes[i]);
}
nodes.push_back(std::move(node));
+1
View File
@@ -33,6 +33,7 @@ ggml_metal_event_t ggml_metal_get_ev_cpy(ggml_metal_t ctx);
void ggml_metal_set_n_cb (ggml_metal_t ctx, int n_cb);
void ggml_metal_set_abort_callback (ggml_metal_t ctx, ggml_abort_callback abort_callback, void * user_data);
bool ggml_metal_supports_family (ggml_metal_t ctx, int family);
void ggml_metal_capture_next_compute(ggml_metal_t ctx);
+33 -22
View File
@@ -6,6 +6,7 @@
#import "ggml-metal-impl.h"
#import "ggml-metal-common.h"
#import "ggml-metal-ops.h"
#import "ggml-metal-fusion.h"
#import <Foundation/Foundation.h>
@@ -36,15 +37,12 @@ struct ggml_metal {
// additional, inference-time compiled pipelines
ggml_metal_pipelines_t pipelines_ext;
bool use_fusion;
bool use_concurrency;
bool use_graph_optimize;
int debug_graph;
int debug_fusion;
// how many times a given op was fused
uint64_t fuse_cnt[GGML_OP_COUNT];
struct ggml_metal_fusion_info * finfo;
// capture state
int capture_compute;
@@ -139,7 +137,6 @@ ggml_metal_t ggml_metal_init(ggml_metal_device_t dev) {
res->d_queue = dispatch_queue_create("ggml-metal", DISPATCH_QUEUE_CONCURRENT);
res->use_fusion = getenv("GGML_METAL_FUSION_DISABLE") == nil;
res->use_concurrency = getenv("GGML_METAL_CONCURRENCY_DISABLE") == nil;
{
@@ -147,20 +144,19 @@ ggml_metal_t ggml_metal_init(ggml_metal_device_t dev) {
res->debug_graph = val ? atoi(val) : 0;
}
{
const char * val = getenv("GGML_METAL_FUSION_DEBUG");
res->debug_fusion = val ? atoi(val) : 0;
}
res->use_graph_optimize = true;
if (getenv("GGML_METAL_GRAPH_OPTIMIZE_DISABLE") != NULL) {
res->use_graph_optimize = false;
}
memset(res->fuse_cnt, 0, sizeof(res->fuse_cnt));
res->finfo = ggml_metal_device_get_fusion_info(dev);
if (ggml_metal_fusion_info_stats(res->finfo)) {
ggml_metal_fusion_info_labels_init(res->finfo);
res->n_cb = 0;
}
GGML_LOG_INFO("%s: use fusion = %s\n", __func__, res->use_fusion ? "true" : "false");
GGML_LOG_INFO("%s: use fusion = %s\n", __func__, ggml_metal_fusion_info_enabled(res->finfo) ? "true" : "false");
GGML_LOG_INFO("%s: use concurrency = %s\n", __func__, res->use_concurrency ? "true" : "false");
GGML_LOG_INFO("%s: use graph optimize = %s\n", __func__, res->use_graph_optimize ? "true" : "false");
@@ -222,15 +218,18 @@ void ggml_metal_free(ggml_metal_t ctx) {
ctx->pipelines_ext = nil;
}
if (ctx->debug_fusion > 0) {
if (ggml_metal_fusion_info_debug(ctx->finfo) > 0) {
GGML_LOG_DEBUG("%s: fusion stats:\n", __func__);
for (int i = 0; i < GGML_OP_COUNT; i++) {
if (ctx->fuse_cnt[i] == 0) {
const int n_fusions = ggml_metal_fusion_info_n_fusions(ctx->finfo);
for (int i = 0; i < n_fusions; i++) {
const uint64_t count = ggml_metal_fusion_info_count(ctx->finfo, i);
if (count == 0) {
continue;
}
// note: cannot use ggml_log here
GGML_LOG_DEBUG("%s: - %s: %" PRIu64 "\n", __func__, ggml_op_name((enum ggml_op) i), ctx->fuse_cnt[i]);
GGML_LOG_DEBUG("%s: - %s: %" PRIu64 "\n", __func__, ggml_metal_fusion_info_label(ctx->finfo, i), count);
}
}
@@ -481,10 +480,17 @@ enum ggml_status ggml_metal_graph_compute(ggml_metal_t ctx, struct ggml_cgraph *
@autoreleasepool {
ctx->gf = gf;
ctx->n_nodes_0 = MIN(n_main, gf->n_nodes);
ctx->n_nodes_1 = gf->n_nodes - ctx->n_nodes_0;
if (ctx->n_cb == 0) {
// single-threaded encoding: the whole graph is encoded by one command buffer
ctx->n_nodes_0 = gf->n_nodes;
ctx->n_nodes_1 = 0;
ctx->n_nodes_per_cb = 0;
} else {
ctx->n_nodes_0 = MIN(n_main, gf->n_nodes);
ctx->n_nodes_1 = gf->n_nodes - ctx->n_nodes_0;
ctx->n_nodes_per_cb = (ctx->n_nodes_1 + ctx->n_cb - 1) / ctx->n_cb;
ctx->n_nodes_per_cb = (ctx->n_nodes_1 + ctx->n_cb - 1) / ctx->n_cb;
}
if (ctx->capture_compute >= 0) {
ctx->capture_compute--;
@@ -682,6 +688,12 @@ ggml_metal_event_t ggml_metal_get_ev_cpy(ggml_metal_t ctx) {
}
void ggml_metal_set_n_cb(ggml_metal_t ctx, int n_cb) {
// when fusion stats are collected the graph must be encoded by a single thread so the
// counters are race-free; override whatever the caller requested
if (ggml_metal_fusion_info_stats(ctx->finfo)) {
n_cb = 0;
}
if (ctx->n_cb != n_cb) {
ctx->n_cb = MIN(n_cb, GGML_METAL_MAX_COMMAND_BUFFERS);
@@ -717,13 +729,12 @@ void ggml_metal_set_n_cb(ggml_metal_t ctx, int n_cb) {
ctx->dev,
cmd_buf,
ctx->gf,
ctx->finfo,
idx_start,
idx_end,
ctx->use_fusion,
ctx->use_concurrency,
ctx->capture_compute,
ctx->debug_graph,
ctx->debug_fusion);
ctx->debug_graph);
for (int idx = 0; idx < ggml_metal_op_n_nodes(ctx_op); ++idx) {
const int res = ggml_metal_op_encode(ctx_op, idx);
+72
View File
@@ -932,12 +932,24 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
nsg = N_SG_IQ2_XXS;
nr0 = N_R0_IQ2_XXS;
smem = 256*8+128;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ2_XXS_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ2_XS:
{
nsg = N_SG_IQ2_XS;
nr0 = N_R0_IQ2_XS;
smem = 512*8+128;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ2_XS_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ3_XXS:
{
@@ -957,21 +969,45 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
nsg = N_SG_IQ3_S;
nr0 = N_R0_IQ3_S;
smem = 512*4;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ3_S_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ2_S:
{
nsg = N_SG_IQ2_S;
nr0 = N_R0_IQ2_S;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ2_S_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ1_S:
{
nsg = N_SG_IQ1_S;
nr0 = N_R0_IQ1_S;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ1_S_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ1_M:
{
nsg = N_SG_IQ1_M;
nr0 = N_R0_IQ1_M;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ1_M_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ4_NL:
{
@@ -1177,12 +1213,24 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id(ggml_m
nsg = N_SG_IQ2_XXS;
nr0 = N_R0_IQ2_XXS;
smem = 256*8+128;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ2_XXS_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ2_XS:
{
nsg = N_SG_IQ2_XS;
nr0 = N_R0_IQ2_XS;
smem = 512*8+128;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ2_XS_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ3_XXS:
{
@@ -1202,21 +1250,45 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id(ggml_m
nsg = N_SG_IQ3_S;
nr0 = N_R0_IQ3_S;
smem = 512*4;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ3_S_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ2_S:
{
nsg = N_SG_IQ2_S;
nr0 = N_R0_IQ2_S;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ2_S_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ1_S:
{
nsg = N_SG_IQ1_S;
nr0 = N_R0_IQ1_S;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ1_S_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ1_M:
{
nsg = N_SG_IQ1_M;
nr0 = N_R0_IQ1_M;
const int nb32 = ne00/32;
if (nb32 < 32 && (32 % nb32) == 0) {
nr0 = N_R0_IQ1_M_SPLIT;
split = true;
}
} break;
case GGML_TYPE_IQ4_NL:
{
+5
View File
@@ -325,6 +325,11 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
const struct ggml_metal_device_props * ggml_metal_device_get_props(ggml_metal_device_t dev);
struct ggml_metal_fusion_info;
// the device-owned fusion debugging context (NULL unless fusion debugging is enabled)
struct ggml_metal_fusion_info * ggml_metal_device_get_fusion_info(ggml_metal_device_t dev);
//
// device buffers
//
+17
View File
@@ -1,4 +1,5 @@
#import "ggml-metal-device.h"
#import "ggml-metal-fusion.h"
#import "ggml-impl.h"
#import "ggml-backend-impl.h"
@@ -896,6 +897,9 @@ struct ggml_metal_device {
struct ggml_metal_device_props props;
// shared fusion debugging context
struct ggml_metal_fusion_info * finfo;
// virtual address for GPU memory allocations
atomic_uintptr_t addr_virt;
};
@@ -1274,6 +1278,13 @@ ggml_metal_device_t ggml_metal_device_init(int device, int n_devices) {
dev->props.max_working_set_size = dev->mtl_device.maxBufferLength;
}
{
const char * val = getenv("GGML_METAL_FUSION_DEBUG");
dev->finfo = ggml_metal_fusion_info_init(
getenv("GGML_METAL_FUSION_DISABLE") == nil,
val ? atoi(val) : 0);
}
snprintf(dev->props.name, sizeof(dev->props.name), "%s%d", "MTL", device);
const char * gpu_name = [[dev->mtl_device name] UTF8String];
if (n_devices > 1) {
@@ -1348,6 +1359,8 @@ void ggml_metal_device_free(ggml_metal_device_t dev) {
assert(dev != NULL);
@autoreleasepool {
ggml_metal_fusion_info_free(dev->finfo);
ggml_metal_rsets_free(dev->rsets);
ggml_metal_library_free(dev->library);
@@ -1935,6 +1948,10 @@ static void ggml_metal_device_disable_tensor(ggml_metal_device_t dev) {
dev->props.has_tensor = false;
}
struct ggml_metal_fusion_info * ggml_metal_device_get_fusion_info(ggml_metal_device_t dev) {
return dev->finfo;
}
//
// device buffers
//
+502
View File
@@ -0,0 +1,502 @@
#include "ggml-metal-fusion.h"
#include "ggml-backend-impl.h"
#include "ggml-metal-device.h"
#include <algorithm>
#include <string>
#include <vector>
// ---- helpers -------------------------------------------------------------
// true if two tensors live in the same Metal buffer
static bool ggml_metal_fusion_same_buffer(const ggml_tensor * a, const ggml_tensor * b) {
if (!a || !b) {
return false;
}
ggml_backend_buffer_t ba = a->view_src ? a->view_src->buffer : a->buffer;
ggml_backend_buffer_t bb = b->view_src ? b->view_src->buffer : b->buffer;
ggml_metal_buffer_t ca = (ggml_metal_buffer_t) ba->context;
ggml_metal_buffer_t cb = (ggml_metal_buffer_t) bb->context;
return ggml_metal_buffer_get_id(ca, a).metal == ggml_metal_buffer_get_id(cb, b).metal;
}
// ---- pattern checks ------------------------------------------------------
// NORM/RMS_NORM + MUL + ADD: the weight/bias of each fused step must match the norm input
// width, be contiguous rows, and the fused outputs must stay F32
static bool ggml_metal_fusion_check_norm(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
ggml_metal_fusion_mode mode) {
GGML_UNUSED(mode);
GGML_ASSERT(fusion->n_ops >= 2);
for (int j = 1; j < fusion->n_ops; j++) {
// the fused MUL/ADD must read the previous node as src0
if (nodes[j]->src[0] != nodes[j - 1]) {
return false;
}
// the weight/bias must have the same row width as the norm input
if (nodes[j]->src[1]->ne[0] != nodes[0]->ne[0]) {
return false;
}
if (!ggml_is_contiguous_rows(nodes[j]->src[1])) {
return false;
}
if (nodes[j]->type != GGML_TYPE_F32) {
return false;
}
}
return true;
}
// ADD x N: each ADD reads the previous ADD as src0, and all addends must share layout
// (and, in FULL mode, live in the same Metal buffer)
static bool ggml_metal_fusion_check_add_chain(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
ggml_metal_fusion_mode mode) {
GGML_ASSERT(fusion->n_ops >= 2);
for (int j = 1; j < fusion->n_ops; j++) {
if (nodes[j]->src[0] != nodes[j - 1]) {
return false;
}
if (!ggml_are_same_layout(nodes[j]->src[1], nodes[j - 1]->src[1])) {
return false;
}
if (mode == GGML_METAL_FUSION_FULL) {
if (!ggml_metal_fusion_same_buffer(nodes[j]->src[1], nodes[0]->src[1])) {
return false;
}
}
}
return true;
}
// GATED_DELTA_NET + CPY: the trailing cpy scatters the gdn state snapshots into the recurrent
// cache, so the gdn kernel writes them straight to the cache and the cpy is elided.
// mirrors ggml_metal_op_can_fuse_gdn_cache (PR #25788). the gdn output has other consumers (the
// attn scores view), so unlike the other patterns this is not an elision chain: the structural
// checks live entirely in this callback (unsafe = true).
static bool ggml_metal_fusion_check_gdn_cache(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
ggml_metal_fusion_mode mode) {
GGML_UNUSED(fusion);
const ggml_tensor * gdn = nodes[0];
const ggml_tensor * cpy = nodes[1];
// the kernel skips the snapshot tail, so the gdn output must not be a graph output
if (gdn->type != GGML_TYPE_F32 || (gdn->flags & GGML_TENSOR_FLAG_OUTPUT)) {
return false;
}
if (cpy->op != GGML_OP_CPY || (cpy->flags & GGML_TENSOR_FLAG_OUTPUT)) {
return false;
}
const int64_t S_v = gdn->src[2]->ne[0];
const int64_t H = gdn->src[2]->ne[1];
const int64_t n_tokens = gdn->src[2]->ne[2];
const int64_t n_seqs = gdn->src[2]->ne[3];
const int64_t K = ggml_get_op_params_i32(gdn, 0);
const size_t tail_off = ggml_row_size(GGML_TYPE_F32, S_v * H * n_tokens * n_seqs);
const int64_t D = S_v * S_v * H;
const int64_t n_written = std::min<int64_t>(n_tokens, K);
const ggml_tensor * src = cpy->src[0]; // gdn snapshot tail view
const ggml_tensor * dst = cpy->src[1]; // cache view
// src must be this gdn's snapshot tail (contiguous, at the tail offset)
if (src->op != GGML_OP_VIEW || src->view_src != gdn ||
src->view_offs != tail_off || !ggml_is_contiguous(src)) {
return false;
}
const int64_t expected_ne[GGML_MAX_DIMS] = { D, n_seqs, n_written, 1 };
if (dst->type != GGML_TYPE_F32 ||
!std::equal(expected_ne, expected_ne + GGML_MAX_DIMS, dst->ne) ||
dst->nb[0] != ggml_type_size(GGML_TYPE_F32) ||
dst->nb[1] != ggml_row_size(GGML_TYPE_F32, D)) {
return false;
}
if (mode == GGML_METAL_FUSION_FULL) {
// the cache must be allocated so the kernel can write straight to its buffer
if (dst->data == nullptr) {
return false;
}
}
return true;
}
// MUL + SIN + SQR + MUL + ADD (snake activation)
static bool ggml_metal_fusion_check_snake(
const ggml_metal_fusion * fusion,
const ggml_tensor * const * nodes,
ggml_metal_fusion_mode mode) {
GGML_UNUSED(fusion);
GGML_UNUSED(mode);
const ggml_tensor * mul0 = nodes[0];
const ggml_tensor * sin_node = nodes[1];
const ggml_tensor * sqr = nodes[2];
const ggml_tensor * mul1 = nodes[3];
const ggml_tensor * add = nodes[4];
// x carries the full activation shape, a is the broadcast operand
const ggml_tensor * x = ggml_are_same_shape(mul0, mul0->src[0]) ? mul0->src[0] : mul0->src[1];
const ggml_tensor * a = (x == mul0->src[0]) ? mul0->src[1] : mul0->src[0];
// mul1 reads sqr and inv_b in either operand order
const ggml_tensor * inv_b = (mul1->src[0] == sqr) ? mul1->src[1] : mul1->src[0];
// closure check: the trailing add reads the same x as the leading mul
const ggml_tensor * x_in_add = (add->src[0] == mul1) ? add->src[1] : add->src[0];
// x is in the supported whitelist and every chain intermediate shares x's type.
// a and inv_b bind as device const float * in the kernel, so they stay F32.
const bool types_ok =
(x->type == GGML_TYPE_F32 || x->type == GGML_TYPE_F16 || x->type == GGML_TYPE_BF16) &&
(a->type == GGML_TYPE_F32) && (inv_b->type == GGML_TYPE_F32) &&
(mul0->type == x->type) && (sin_node->type == x->type) &&
(sqr->type == x->type) && (mul1->type == x->type) &&
(add->type == x->type);
// a / inv_b collapse to [1, C, 1, 1], x and add stay 2D
const bool shape_ok = ggml_are_same_shape(a, inv_b) && a->ne[0] == 1 && a->ne[1] == x->ne[1];
const bool dim_ok =
(x->ne[2] == 1) && (x->ne[3] == 1) &&
(add->ne[2] == 1) && (add->ne[3] == 1) &&
(a->ne[2] == 1) && (a->ne[3] == 1) &&
(inv_b->ne[2] == 1) && (inv_b->ne[3] == 1);
// kernel reads x[idx] and a[c] / inv_b[c] linearly, so every operand is contiguous
const bool contig_ok =
ggml_is_contiguous(x) && ggml_is_contiguous(add) &&
ggml_is_contiguous(a) && ggml_is_contiguous(inv_b);
return types_ok && shape_ok && dim_ok && contig_ok && x_in_add == x;
}
// ---- patterns ------------------------------------------------------------
static const ggml_op ops_norm_mul[] = { GGML_OP_NORM, GGML_OP_MUL };
static const ggml_op ops_norm_mul_add[] = { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD };
static const ggml_op ops_rms_norm_mul[] = { GGML_OP_RMS_NORM, GGML_OP_MUL };
static const ggml_op ops_rms_norm_mul_add[] = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD };
static const ggml_op ops_add_2[] = { GGML_OP_ADD, GGML_OP_ADD };
static const ggml_op ops_add_3[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const ggml_op ops_add_4[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const ggml_op ops_add_5[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const ggml_op ops_add_6[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const ggml_op ops_add_7[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const ggml_op ops_snake[] = { GGML_OP_MUL, GGML_OP_SIN, GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD };
static const ggml_op ops_gdn_cache[] = { GGML_OP_GATED_DELTA_NET, GGML_OP_CPY };
static const ggml_metal_fusion ggml_metal_fusions[] = {
{ GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, 2, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, 3, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, 2, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, 3, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, 2, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, 3, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, 4, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, 5, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, 6, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, 7, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_SNAKE, ops_snake, 5, false, ggml_metal_fusion_check_snake },
{ GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, 2, true, ggml_metal_fusion_check_gdn_cache },
};
const ggml_metal_fusion * ggml_metal_fusion_all(int * n) {
*n = (int) sizeof(ggml_metal_fusions) / sizeof(ggml_metal_fusions[0]);
return ggml_metal_fusions;
}
// ---- shared fusion info ---------------------------------------------------
static std::string ggml_metal_fusion_label(const ggml_metal_fusion * fusion) {
GGML_ASSERT(fusion != nullptr);
std::string label;
for (int j = 0; j < fusion->n_ops; j++) {
if (j > 0) {
label += '+';
}
label += ggml_op_name(fusion->ops[j]);
}
return label;
}
struct ggml_metal_fusion_info {
std::vector<std::string> labels;
std::vector<uint64_t> counts;
bool enabled;
bool stats;
bool labels_set;
int debug;
};
struct ggml_metal_fusion_info * ggml_metal_fusion_info_init(bool enabled, int debug) {
ggml_metal_fusion_info * finfo = new ggml_metal_fusion_info;
finfo->enabled = enabled;
finfo->stats = debug > 0;
finfo->labels_set = false;
finfo->debug = debug;
if (finfo->stats) {
ggml_metal_fusion_info_labels_init(finfo);
}
return finfo;
}
void ggml_metal_fusion_info_free(struct ggml_metal_fusion_info * finfo) {
delete finfo;
}
bool ggml_metal_fusion_info_enabled(const struct ggml_metal_fusion_info * finfo) {
return finfo->enabled;
}
bool ggml_metal_fusion_info_stats(const struct ggml_metal_fusion_info * finfo) {
return finfo->stats;
}
int ggml_metal_fusion_info_debug(const struct ggml_metal_fusion_info * finfo) {
return finfo->debug;
}
int ggml_metal_fusion_info_n_fusions(const struct ggml_metal_fusion_info * finfo) {
return (int) finfo->labels.size();
}
const char * ggml_metal_fusion_info_label(const struct ggml_metal_fusion_info * finfo, int idx) {
GGML_ASSERT(idx >= 0 && idx < (int) finfo->labels.size());
return finfo->labels[idx].c_str();
}
uint64_t ggml_metal_fusion_info_count(const struct ggml_metal_fusion_info * finfo, int idx) {
GGML_ASSERT(idx >= 0 && idx < (int) finfo->counts.size());
return finfo->counts[idx];
}
void ggml_metal_fusion_info_count_fusion(struct ggml_metal_fusion_info * finfo, const struct ggml_metal_fusion * fusion) {
if (!finfo->stats || fusion == nullptr) {
return;
}
int n = 0;
const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
int idx = -1;
for (int i = 0; i < n; i++) {
if (&all[i] == fusion) {
idx = i;
break;
}
}
if (idx >= 0 && idx < (int) finfo->counts.size()) {
finfo->counts[idx]++;
}
}
void ggml_metal_fusion_info_set_enabled(struct ggml_metal_fusion_info * finfo, bool enabled) {
finfo->enabled = enabled;
}
void ggml_metal_fusion_info_labels_init(struct ggml_metal_fusion_info * finfo) {
if (finfo->labels_set) {
return;
}
int n = 0;
const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
finfo->labels.clear();
finfo->counts.assign(n, 0);
finfo->labels.reserve(n);
for (int i = 0; i < n; i++) {
finfo->labels.emplace_back(ggml_metal_fusion_label(&all[i]));
}
finfo->labels_set = true;
}
void ggml_metal_fusion_info_stats_init(struct ggml_metal_fusion_info * finfo) {
finfo->stats = true;
ggml_metal_fusion_info_labels_init(finfo);
}
void ggml_metal_fusion_info_stats_reset(struct ggml_metal_fusion_info * finfo) {
std::fill(finfo->counts.begin(), finfo->counts.end(), 0);
}
int ggml_metal_fusion_info_stats_get(const struct ggml_metal_fusion_info * finfo, const char ** labels, uint64_t * counts, int n) {
const int n_fusions = (int) finfo->labels.size();
if (labels == nullptr) {
return n_fusions;
}
const int n_fill = std::min(n, n_fusions);
for (int i = 0; i < n_fill; i++) {
labels[i] = finfo->labels[i].c_str();
if (counts != nullptr) {
counts[i] = finfo->counts[i];
}
}
return n_fill;
}
// ---- queries -------------------------------------------------------------
// find the longest pattern matching the node sequence starting at idx
// (idx is a position in node_idxs, which maps to graph node indices)
const ggml_metal_fusion * ggml_metal_fusion_next(
const ggml_cgraph * gf,
const int * node_idxs,
int n_idxs,
int idx,
ggml_metal_fusion_mode mode,
int * n_out) {
int n = 0;
const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
const ggml_metal_fusion * res = nullptr;
int best = 1;
for (int i = 0; i < n; i++) {
const ggml_metal_fusion * fusion = &all[i];
// only look for a longer match than the current best
if (fusion->n_ops <= best) {
continue;
}
if (idx + fusion->n_ops > n_idxs) {
continue;
}
const ggml_tensor * nodes[GGML_METAL_FUSION_MAX];
// the op sequence must match exactly
bool ok = true;
for (int j = 0; j < fusion->n_ops; j++) {
nodes[j] = gf->nodes[node_idxs[idx + j]];
if (nodes[j]->op != fusion->ops[j]) {
ok = false;
break;
}
}
if (!ok) {
continue;
}
if (!fusion->unsafe) {
// common element-wise chain constraints: each node reads the previous one,
// and all nodes have the same shape
for (int j = 1; j < fusion->n_ops && ok; j++) {
if (nodes[j]->src[0] != nodes[j - 1] && nodes[j]->src[1] != nodes[j - 1]) {
ok = false;
break;
}
if (!ggml_are_same_shape(nodes[j], nodes[j - 1])) {
ok = false;
break;
}
}
if (!ok) {
continue;
}
// all current fusions are single-output elision chains, so the last node is the only output
// TODO: multi-output fusions: store pattern-relative offsets in the table and translate them here
int outputs_buf[1];
outputs_buf[0] = node_idxs[idx + fusion->n_ops - 1];
// structural subgraph checks (op sequence, elidable uses, view containment)
if (!ggml_can_fuse_subgraph_ext(gf, node_idxs + idx, fusion->n_ops, fusion->ops, outputs_buf, 1)) {
continue;
}
}
// pattern-specific checks (the sole validator for unsafe patterns)
if (fusion->check && !fusion->check(fusion, nodes, mode)) {
continue;
}
best = fusion->n_ops;
res = fusion;
}
*n_out = best;
return res;
}
// optimize phase: maximum number of nodes starting at idx (a raw sequential graph index) that
// could be fused, chaining patterns back-to-back. matching runs on the same filtered (view
// transparent) node sequence that the compute phase uses, so the returned count is the raw index
// span from idx to the last matched node (intermediate views are packed along).
int ggml_metal_fusion_max(const ggml_cgraph * gf, int idx) {
// an empty/view node cannot start a pattern - pack it alone
if (ggml_op_is_empty(gf->nodes[idx]->op) || ggml_is_empty(gf->nodes[idx])) {
return 1;
}
// collect the non-empty node indices starting at idx
int idxs[GGML_METAL_FUSION_MAX];
int n_idxs = 0;
for (int i = idx; i < gf->n_nodes && n_idxs < GGML_METAL_FUSION_MAX; i++) {
if (!ggml_op_is_empty(gf->nodes[i]->op) && !ggml_is_empty(gf->nodes[i])) {
idxs[n_idxs++] = i;
}
}
if (n_idxs == 0) {
return 1;
}
int total = 0;
int i_f = 0;
while (i_f < n_idxs && total < GGML_METAL_FUSION_MAX) {
int len = 1;
const ggml_metal_fusion * fusion = ggml_metal_fusion_next(gf, idxs, n_idxs, i_f, GGML_METAL_FUSION_STRUCTURAL, &len);
if (!fusion || total + len > GGML_METAL_FUSION_MAX) {
break;
}
total += len;
i_f += len;
}
if (i_f == 0) {
return 1;
}
// map the matched non-empty nodes back to the raw index span (views are included)
return std::min(GGML_METAL_FUSION_MAX, idxs[i_f - 1] - idx + 1);
}
+104
View File
@@ -0,0 +1,104 @@
// single source of truth for the fusions supported by the Metal backend
//
// every fusable subgraph is declared exactly once as a ggml_metal_fusion entry in
// the table in ggml-metal-fusion.cpp. both the graph optimizer (ggml_metal_fusion_max)
// and the op encoders (ggml_metal_fusion_next) consult this same table, so the two
// phases can never disagree about what can be fused.
#pragma once
#include "ggml-impl.h"
#include <stdbool.h>
#ifdef __cplusplus
extern "C" {
#endif
// the maximum number of nodes that can be fused in a single kernel
// (also the maximum length of a packed fusion group during graph optimization)
#define GGML_METAL_FUSION_MAX 16
typedef enum ggml_metal_fusion_mode {
// structural checks only; used by the graph optimizer, at which point the graph
// tensors are not allocated yet, so buffer placement cannot be verified
GGML_METAL_FUSION_STRUCTURAL = 0,
// full checks, including buffer placement; used by the op encoders
GGML_METAL_FUSION_FULL,
} ggml_metal_fusion_mode;
// identifier of each fusion pattern so the op encoders know which kernel to use
typedef enum ggml_metal_fusion_id {
GGML_METAL_FUSION_NONE = 0,
GGML_METAL_FUSION_NORM_MUL, // NORM/RMS_NORM + MUL
GGML_METAL_FUSION_NORM_MUL_ADD, // NORM/RMS_NORM + MUL + ADD
GGML_METAL_FUSION_ADD_CHAIN, // ADD x N (N in [2, 7])
GGML_METAL_FUSION_SNAKE, // MUL + SIN + SQR + MUL + ADD
GGML_METAL_FUSION_GDN_CACHE, // GATED_DELTA_NET + CPY (write snapshots into the recurrent cache)
} ggml_metal_fusion_id;
struct ggml_metal_fusion {
ggml_metal_fusion_id id;
const enum ggml_op * ops; // op sequence (fixed length)
int n_ops; // number of ops
// if unsafe: the generic chain/shape + ggml_can_fuse_subgraph checks are skipped and the
// check callback below is the sole validator (used for patterns that are not elision chains,
// e.g. the gdn + cache-cpy write-through fusion)
bool unsafe;
// extra backend constraints on top of ggml_can_fuse_subgraph
// nodes[j] is the j-th node of the pattern
bool (*check)(const struct ggml_metal_fusion * fusion,
const struct ggml_tensor * const * nodes,
ggml_metal_fusion_mode mode);
};
typedef struct ggml_metal_fusion ggml_metal_fusion;
// the single table of all fusions supported by the Metal backend
const ggml_metal_fusion * ggml_metal_fusion_all(int * n);
// ---- shared fusion info ---------------------------------------------------
// shared fusion debugging context, owned by the device; newly created backend contexts for that
// device register with it so the fusion counters are race-free and accumulate across contexts.
struct ggml_metal_fusion_info; // defined in ggml-metal-fusion.cpp
struct ggml_metal_fusion_info * ggml_metal_fusion_info_init(bool enabled, int debug);
void ggml_metal_fusion_info_free(struct ggml_metal_fusion_info * finfo);
bool ggml_metal_fusion_info_enabled(const struct ggml_metal_fusion_info * finfo);
bool ggml_metal_fusion_info_stats (const struct ggml_metal_fusion_info * finfo);
int ggml_metal_fusion_info_debug (const struct ggml_metal_fusion_info * finfo);
int ggml_metal_fusion_info_n_fusions(const struct ggml_metal_fusion_info * finfo);
const char * ggml_metal_fusion_info_label (const struct ggml_metal_fusion_info * finfo, int idx);
uint64_t ggml_metal_fusion_info_count (const struct ggml_metal_fusion_info * finfo, int idx);
void ggml_metal_fusion_info_count_fusion(struct ggml_metal_fusion_info * finfo, const struct ggml_metal_fusion * fusion);
void ggml_metal_fusion_info_set_enabled (struct ggml_metal_fusion_info * finfo, bool enabled);
void ggml_metal_fusion_info_stats_init ( struct ggml_metal_fusion_info * finfo);
void ggml_metal_fusion_info_stats_reset( struct ggml_metal_fusion_info * finfo);
int ggml_metal_fusion_info_stats_get (const struct ggml_metal_fusion_info * finfo, const char ** labels, uint64_t * counts, int n);
void ggml_metal_fusion_info_labels_init( struct ggml_metal_fusion_info * finfo);
// compute phase: longest fusion starting at idx (a position in node_idxs) that matches in `mode`.
// returns the matching pattern (nullptr if no fusion) and sets *n_out to the number of nodes consumed.
const ggml_metal_fusion * ggml_metal_fusion_next(
const struct ggml_cgraph * gf,
const int * node_idxs,
int n_idxs,
int idx,
ggml_metal_fusion_mode mode,
int * n_out);
// optimize phase: maximum number of nodes starting at idx (a raw sequential graph index) that
// could be fused, chaining patterns back-to-back. returns at least 1.
int ggml_metal_fusion_max(const struct ggml_cgraph * gf, int idx);
#ifdef __cplusplus
}
#endif
+7
View File
@@ -62,18 +62,23 @@
#define N_R0_IQ1_S 4
#define N_SG_IQ1_S 2
#define N_R0_IQ1_S_SPLIT 8
#define N_R0_IQ1_M 4
#define N_SG_IQ1_M 2
#define N_R0_IQ1_M_SPLIT 8
#define N_R0_IQ2_XXS 4
#define N_SG_IQ2_XXS 2
#define N_R0_IQ2_XXS_SPLIT 8
#define N_R0_IQ2_XS 4
#define N_SG_IQ2_XS 2
#define N_R0_IQ2_XS_SPLIT 8
#define N_R0_IQ2_S 4
#define N_SG_IQ2_S 2
#define N_R0_IQ2_S_SPLIT 8
#define N_R0_IQ3_XXS 4
#define N_SG_IQ3_XXS 2
@@ -81,6 +86,7 @@
#define N_R0_IQ3_S 4
#define N_SG_IQ3_S 2
#define N_R0_IQ3_S_SPLIT 8
#define N_R0_IQ4_NL 2
#define N_SG_IQ4_NL 2
@@ -979,6 +985,7 @@ typedef struct {
uint64_t nb1;
uint64_t nb2;
uint64_t nb3;
uint64_t nb_out; // 0 => snapshots are appended after the attn scores (unfused)
} ggml_metal_kargs_gated_delta_net;
typedef struct {
+107 -167
View File
@@ -7,6 +7,7 @@
#include "ggml-metal-impl.h"
#include "ggml-metal-common.h"
#include "ggml-metal-device.h"
#include "ggml-metal-fusion.h"
#include "ggml-metal-tuning.h"
#include <cassert>
@@ -31,24 +32,22 @@ struct ggml_metal_op {
ggml_metal_device_t dev,
ggml_metal_cmd_buf_t cmd_buf,
ggml_cgraph * gf,
ggml_metal_fusion_info * finfo,
int idx_start,
int idx_end,
bool use_fusion,
bool use_concurrency,
bool use_capture,
int debug_graph,
int debug_fusion) {
int debug_graph) {
this->dev = dev;
this->lib = ggml_metal_device_get_library(dev);
this->enc = ggml_metal_encoder_init(cmd_buf, use_concurrency);
this->mem_ranges = ggml_mem_ranges_init(debug_graph);
this->finfo = finfo;
this->idx_start = idx_start;
this->idx_end = idx_end;
this->use_fusion = use_fusion;
this->use_concurrency = use_concurrency;
this->use_capture = use_capture;
this->debug_graph = debug_graph;
this->debug_fusion = debug_fusion;
this->gf = gf;
idxs.reserve(gf->n_nodes);
@@ -78,15 +77,24 @@ struct ggml_metal_op {
return ggml_graph_node(gf, idxs[i]);
}
bool can_fuse(int i0, const ggml_op * ops, int n_ops) const {
assert(use_fusion);
// consult the fusion table for the longest pattern starting at i0
// returns the matching pattern (nullptr if no fusion) and sets *n_out to the number of nodes
const ggml_metal_fusion * can_fuse(int i0, enum ggml_metal_fusion_mode mode, int * n_out) const {
assert(use_fusion());
assert(i0 >= 0 && i0 < n_nodes());
if (i0 + n_ops > n_nodes()) {
return false;
}
return ggml_metal_fusion_next(gf, idxs.data(), (int) idxs.size(), i0, mode, n_out);
}
return ggml_can_fuse_ext(gf, idxs.data() + i0, ops, n_ops);
// whether to attempt fusion; the toggle lives in the shared fusion debugging context owned
// by the device (initialized from GGML_METAL_FUSION_DISABLE, overridable by the test)
bool use_fusion() const {
return ggml_metal_fusion_info_enabled(finfo);
}
// record that a fusion fired, indexed by the matching table entry
void count_fusions(const ggml_metal_fusion * fusion) const {
ggml_metal_fusion_info_count_fusion(finfo, fusion);
}
ggml_metal_device_t dev;
@@ -94,12 +102,13 @@ struct ggml_metal_op {
ggml_metal_encoder_t enc;
ggml_mem_ranges_t mem_ranges;
bool use_fusion;
// shared fusion debugging context
ggml_metal_fusion_info * finfo;
bool use_concurrency;
bool use_capture;
int debug_graph;
int debug_fusion;
private:
ggml_cgraph * gf;
@@ -115,24 +124,22 @@ ggml_metal_op_t ggml_metal_op_init(
ggml_metal_device_t dev,
ggml_metal_cmd_buf_t cmd_buf,
ggml_cgraph * gf,
ggml_metal_fusion_info * finfo,
int idx_start,
int idx_end,
bool use_fusion,
bool use_concurrency,
bool use_capture,
int debug_graph,
int debug_fusion) {
int debug_graph) {
ggml_metal_op_t res = new ggml_metal_op(
dev,
cmd_buf,
gf,
finfo,
idx_start,
idx_end,
use_fusion,
use_concurrency,
use_capture,
debug_graph,
debug_fusion);
debug_graph);
return res;
}
@@ -1868,6 +1875,8 @@ int ggml_metal_op_gated_delta_net(ggml_metal_op_t ctx, int idx) {
ggml_metal_library_t lib = ctx->lib;
ggml_metal_encoder_t enc = ctx->enc;
const bool use_fusion = ctx->use_fusion();
const int debug_fusion = ggml_metal_fusion_info_debug(ctx->finfo);
GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb);
@@ -1880,6 +1889,31 @@ int ggml_metal_op_gated_delta_net(ggml_metal_op_t ctx, int idx) {
auto pipeline = ggml_metal_library_get_pipeline_gated_delta_net(lib, op);
// when fused with the trailing cache cpy, the snapshots are written straight into the
// recurrent cache and the cpy is skipped (see GGML_METAL_FUSION_GDN_CACHE)
ggml_metal_buffer_id bid_out = ggml_metal_get_buffer_id(op);
uint64_t nb_out = 0;
int n_fuse = 1;
if (use_fusion) {
int n = 1;
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
if (fusion && fusion->id == GGML_METAL_FUSION_GDN_CACHE) {
const ggml_tensor * dst_cache = ctx->node(idx + 1)->src[1]; // cache view
bid_out = ggml_metal_get_buffer_id(dst_cache);
nb_out = dst_cache->nb[2]/sizeof(float);
n_fuse = 2;
ctx->count_fusions(fusion);
if (debug_fusion > 1) {
GGML_LOG_DEBUG("%s: fuse: GATED_DELTA_NET + CPY\n", __func__);
}
}
}
int ida = 0;
ggml_metal_kargs_gated_delta_net args = {
@@ -1918,23 +1952,25 @@ int ggml_metal_op_gated_delta_net(ggml_metal_op_t ctx, int idx) {
/*.nb1 =*/ nb1,
/*.nb2 =*/ nb2,
/*.nb3 =*/ nb3,
/*.nb_out =*/ nb_out,
};
ggml_metal_encoder_set_pipeline(enc, pipeline);
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), ida++);
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), ida++); // args
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), ida++); // q
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), ida++); // k
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[2]), ida++); // v
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[3]), ida++); // gate
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[4]), ida++); // beta
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[5]), ida++); // state
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), ida++); // dst
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), ida++); // dst (attn)
ggml_metal_encoder_set_buffer (enc, bid_out, ida++); // state_out
const int nsg = pipeline.nsg;
ggml_metal_encoder_dispatch_threadgroups(enc, op->src[2]->ne[0]/nsg, op->src[2]->ne[1], op->src[2]->ne[3], 32, nsg, 1);
return 1;
return n_fuse;
}
int ggml_metal_op_solve_tri(ggml_metal_op_t ctx, int idx) {
@@ -3718,56 +3754,20 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) {
return 1;
}
// Snake activation autofuse: mul -> sin -> sqr -> mul -> add
static bool ggml_metal_op_can_fuse_snake(ggml_metal_op_t ctx, int idx) {
static constexpr ggml_op snake_ops[5] = { GGML_OP_MUL, GGML_OP_SIN, GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD };
if (ctx->node(idx)->op != GGML_OP_MUL || !ctx->can_fuse(idx, snake_ops, 5)) {
return false;
}
const ggml_tensor * mul0 = ctx->node(idx + 0);
const ggml_tensor * sin_node = ctx->node(idx + 1);
const ggml_tensor * sqr = ctx->node(idx + 2);
const ggml_tensor * mul1 = ctx->node(idx + 3);
const ggml_tensor * add = ctx->node(idx + 4);
// x carries the full activation shape, a is the broadcast operand
const ggml_tensor * x = ggml_are_same_shape(mul0, mul0->src[0]) ? mul0->src[0] : mul0->src[1];
const ggml_tensor * a = (x == mul0->src[0]) ? mul0->src[1] : mul0->src[0];
// mul1 reads sqr and inv_b in either operand order
const ggml_tensor * inv_b = (mul1->src[0] == sqr) ? mul1->src[1] : mul1->src[0];
// closure check: the trailing add reads the same x as the leading mul
const ggml_tensor * x_in_add = (add->src[0] == mul1) ? add->src[1] : add->src[0];
// x is in the supported whitelist and every chain intermediate shares x's type.
// a and inv_b bind as device const float * in the kernel, so they stay F32.
const bool types_ok =
(x->type == GGML_TYPE_F32 || x->type == GGML_TYPE_F16 || x->type == GGML_TYPE_BF16) &&
(a->type == GGML_TYPE_F32) && (inv_b->type == GGML_TYPE_F32) &&
(mul0->type == x->type) && (sin_node->type == x->type) &&
(sqr->type == x->type) && (mul1->type == x->type) &&
(add->type == x->type);
// a / inv_b collapse to [1, C, 1, 1], x and add stay 2D
const bool shape_ok = ggml_are_same_shape(a, inv_b) && a->ne[0] == 1 && a->ne[1] == x->ne[1];
const bool dim_ok =
(x->ne[2] == 1) && (x->ne[3] == 1) &&
(add->ne[2] == 1) && (add->ne[3] == 1) &&
(a->ne[2] == 1) && (a->ne[3] == 1) &&
(inv_b->ne[2] == 1) && (inv_b->ne[3] == 1);
// kernel reads x[idx] and a[c] / inv_b[c] linearly, so every operand is contiguous
const bool contig_ok =
ggml_is_contiguous(x) && ggml_is_contiguous(add) &&
ggml_is_contiguous(a) && ggml_is_contiguous(inv_b);
return types_ok && shape_ok && dim_ok && contig_ok && x_in_add == x;
}
int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
if (ctx->use_fusion && ggml_metal_op_can_fuse_snake(ctx, idx)) {
return ggml_metal_op_snake_fused(ctx, idx);
int n_fuse = 1;
const ggml_metal_fusion * fusion = nullptr;
if (ctx->use_fusion()) {
int n = 1;
fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
n_fuse = n;
// snake activation autofuse: mul -> sin -> sqr -> mul -> add
if (fusion && fusion->id == GGML_METAL_FUSION_SNAKE) {
ctx->count_fusions(fusion);
return ggml_metal_op_snake_fused(ctx, idx);
}
}
ggml_tensor * op = ctx->node(idx);
@@ -3775,9 +3775,9 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
ggml_metal_library_t lib = ctx->lib;
ggml_metal_encoder_t enc = ctx->enc;
const bool use_fusion = ctx->use_fusion;
const bool use_fusion = ctx->use_fusion();
const int debug_fusion = ctx->debug_fusion;
const int debug_fusion = ggml_metal_fusion_info_debug(ctx->finfo);
GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb);
@@ -3822,57 +3822,19 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
/*.o1 =*/ { bid_src1.offs },
};
ggml_op fops[8];
int n_fuse = 1;
// c[0] = add(a, b[0])
// c[1] = add(c[0], b[1])
// c[2] = add(c[1], b[2])
// ...
if (use_fusion) {
fops[0] = GGML_OP_ADD;
fops[1] = GGML_OP_ADD;
fops[2] = GGML_OP_ADD;
fops[3] = GGML_OP_ADD;
fops[4] = GGML_OP_ADD;
fops[5] = GGML_OP_ADD;
fops[6] = GGML_OP_ADD;
fops[7] = GGML_OP_ADD;
// note: in metal, we sometimes encode the graph in parallel so we have to avoid fusing ops
// across splits. idx_end indicates the last node in the current split
for (n_fuse = 0; n_fuse <= 6; ++n_fuse) {
if (!ctx->can_fuse(idx + n_fuse, fops + n_fuse, 2)) {
break;
}
ggml_tensor * f0 = ctx->node(idx + n_fuse);
ggml_tensor * f1 = ctx->node(idx + n_fuse + 1);
if (f0 != f1->src[0]) {
break;
}
// b[0] === b[1] === ...
if (!ggml_are_same_layout(f0->src[1], f1->src[1])) {
break;
}
// only fuse ops if src1 is in the same Metal buffer
ggml_metal_buffer_id bid_fuse = ggml_metal_get_buffer_id(f1->src[1]);
if (bid_fuse.metal != bid_src1.metal) {
break;
}
//ctx->fuse_cnt[ops[n_fuse + 1]->op]++;
args.o1[n_fuse + 1] = bid_fuse.offs;
if (use_fusion && fusion && fusion->id == GGML_METAL_FUSION_ADD_CHAIN) {
// the offsets of the fused addends are relative to the start of the src1 buffer
for (int i = 1; i < n_fuse; i++) {
args.o1[i] = ggml_metal_get_buffer_id(ctx->node(idx + i)->src[1]).offs;
}
++n_fuse;
ctx->count_fusions(fusion);
if (debug_fusion > 1 && n_fuse > 1) {
if (debug_fusion > 1) {
GGML_LOG_DEBUG("%s: fuse: ADD x %d\n", __func__, n_fuse);
}
}
@@ -4080,9 +4042,9 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
ggml_metal_library_t lib = ctx->lib;
ggml_metal_encoder_t enc = ctx->enc;
const bool use_fusion = ctx->use_fusion;
const bool use_fusion = ctx->use_fusion();
const int debug_fusion = ctx->debug_fusion;
const int debug_fusion = ggml_metal_fusion_info_debug(ctx->finfo);
GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb);
@@ -4110,8 +4072,6 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
/*.nbf3 =*/ { nb03 },
};
ggml_op fops[8];
int n_fuse = 1;
ggml_metal_buffer_id bid_fuse[2] = { bid_src0, bid_src0 };
@@ -4120,55 +4080,35 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
// d[1] = mul(d[0], b)
// d[2] = add(d[1], c)
if (use_fusion) {
fops[0] = op->op;
fops[1] = GGML_OP_MUL;
fops[2] = GGML_OP_ADD;
int n = 1;
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
for (n_fuse = 0; n_fuse <= 1; ++n_fuse) {
if (!ctx->can_fuse(idx + n_fuse, fops + n_fuse, 2)) {
break;
if (fusion && (fusion->id == GGML_METAL_FUSION_NORM_MUL || fusion->id == GGML_METAL_FUSION_NORM_MUL_ADD)) {
n_fuse = n;
ctx->count_fusions(fusion);
for (int i = 1; i < n_fuse; i++) {
const ggml_tensor * fn = ctx->node(idx + i);
bid_fuse[i - 1] = ggml_metal_get_buffer_id(fn->src[1]);
args.nef1[i] = fn->src[1]->ne[1];
args.nef2[i] = fn->src[1]->ne[2];
args.nef3[i] = fn->src[1]->ne[3];
args.nbf1[i] = fn->src[1]->nb[1];
args.nbf2[i] = fn->src[1]->nb[2];
args.nbf3[i] = fn->src[1]->nb[3];
}
ggml_tensor * f0 = ctx->node(idx + n_fuse);
ggml_tensor * f1 = ctx->node(idx + n_fuse + 1);
if (f0 != f1->src[0]) {
break;
}
if (f1->src[1]->ne[0] != op->ne[0]) {
break;
}
if (!ggml_is_contiguous_rows(f1->src[1])) {
break;
}
if (f1->type != GGML_TYPE_F32) {
break;
}
//ctx->fuse_cnt[f1->op]++;
bid_fuse[n_fuse] = ggml_metal_get_buffer_id(f1->src[1]);
args.nef1[n_fuse + 1] = f1->src[1]->ne[1];
args.nef2[n_fuse + 1] = f1->src[1]->ne[2];
args.nef3[n_fuse + 1] = f1->src[1]->ne[3];
args.nbf1[n_fuse + 1] = f1->src[1]->nb[1];
args.nbf2[n_fuse + 1] = f1->src[1]->nb[2];
args.nbf3[n_fuse + 1] = f1->src[1]->nb[3];
}
++n_fuse;
if (debug_fusion > 1 && n_fuse > 1) {
if (n_fuse == 2) {
GGML_LOG_DEBUG("%s: fuse: %s + MUL\n", __func__, ggml_op_name(op->op));
}
if (n_fuse == 3) {
GGML_LOG_DEBUG("%s: fuse: %s + MUL + ADD\n", __func__, ggml_op_name(op->op));
if (debug_fusion > 1) {
if (n_fuse == 2) {
GGML_LOG_DEBUG("%s: fuse: %s + MUL\n", __func__, ggml_op_name(op->op));
}
if (n_fuse == 3) {
GGML_LOG_DEBUG("%s: fuse: %s + MUL + ADD\n", __func__, ggml_op_name(op->op));
}
}
}
}
+4 -3
View File
@@ -8,17 +8,18 @@ extern "C" {
typedef struct ggml_metal_op * ggml_metal_op_t;
struct ggml_metal_fusion; // forward decl (ggml-metal-device.h)
ggml_metal_op_t ggml_metal_op_init(
ggml_metal_device_t dev,
ggml_metal_cmd_buf_t cmd_buf,
struct ggml_cgraph * gf,
struct ggml_metal_fusion_info * finfo,
int idx_start,
int idx_end,
bool use_fusion,
bool use_concurrency,
bool use_capture,
int debug_graph,
int debug_fusion);
int debug_graph);
void ggml_metal_op_free(ggml_metal_op_t ctx);
+42
View File
@@ -4,6 +4,7 @@
#include "ggml-backend-impl.h"
#include "ggml-metal-device.h"
#include "ggml-metal-fusion.h"
#include "ggml-metal-context.h"
#include "ggml-metal-ops.h"
#include "ggml-metal-tuning.h"
@@ -906,6 +907,30 @@ static const char * ggml_backend_metal_tuning_device_token(ggml_backend_dev_t de
return ggml_metal_device_id_token(ggml_metal_device_get_props(ctx_dev)->device_id);
}
// generic fusion debugging API (ad-hoc proc-address mechanism): the test resolves the device
// fusion context once and passes that opaque handle to the rest of the functions
typedef void * ggml_backend_fusion_t;
static ggml_backend_fusion_t ggml_backend_metal_fusion_get(ggml_backend_dev_t dev) {
return ggml_metal_device_get_fusion_info((ggml_metal_device_t)dev->context);
}
static void ggml_backend_metal_fusion_stats_init(ggml_backend_fusion_t finfo) {
ggml_metal_fusion_info_stats_init((struct ggml_metal_fusion_info *) finfo);
}
static void ggml_backend_metal_fusion_stats_reset(ggml_backend_fusion_t finfo) {
ggml_metal_fusion_info_stats_reset((struct ggml_metal_fusion_info *) finfo);
}
static int ggml_backend_metal_fusion_stats_get(ggml_backend_fusion_t finfo, const char ** labels, uint64_t * counts, int n) {
return ggml_metal_fusion_info_stats_get((struct ggml_metal_fusion_info *) finfo, labels, counts, n);
}
static void ggml_backend_metal_fusion_set_enabled(ggml_backend_fusion_t finfo, bool enabled) {
ggml_metal_fusion_info_set_enabled((struct ggml_metal_fusion_info *) finfo, enabled);
}
static void * ggml_backend_metal_get_proc_address(ggml_backend_reg_t reg, const char * name) {
if (strcmp(name, "ggml_backend_get_features") == 0) {
return (void *)ggml_backend_metal_get_features;
@@ -928,6 +953,23 @@ static void * ggml_backend_metal_get_proc_address(ggml_backend_reg_t reg, const
if (strcmp(name, "ggml_backend_metal_tuning_device_token") == 0) {
return (void *)ggml_backend_metal_tuning_device_token;
}
// generic fusion debugging API (ad-hoc proc-address mechanism, not part of the official
// ggml backend interface yet; a backend that adopts it exports these exact names)
if (strcmp(name, "ggml_backend_fusion_get") == 0) {
return (void *)ggml_backend_metal_fusion_get;
}
if (strcmp(name, "ggml_backend_fusion_stats_init") == 0) {
return (void *)ggml_backend_metal_fusion_stats_init;
}
if (strcmp(name, "ggml_backend_fusion_stats_reset") == 0) {
return (void *)ggml_backend_metal_fusion_stats_reset;
}
if (strcmp(name, "ggml_backend_fusion_stats_get") == 0) {
return (void *)ggml_backend_metal_fusion_stats_get;
}
if (strcmp(name, "ggml_backend_fusion_set_enabled") == 0) {
return (void *)ggml_backend_metal_fusion_set_enabled;
}
return NULL;
@@ -15,6 +15,7 @@ kernel void kernel_gated_delta_net_impl(
device const char * b,
device const char * s,
device char * dst,
device char * dst_fuse,
uint3 tgpig[[threadgroup_position_in_grid]],
uint3 tpitg[[thread_position_in_threadgroup]],
uint3 ntg[[threads_per_threadgroup]]) {
@@ -65,6 +66,12 @@ kernel void kernel_gated_delta_net_impl(
// per-(seq,head) offset within a slot
const uint state_out_base = (i23*args.ne21 + i21)*S_v*S_v + i20*S_v;
// when fused with the cache cpy, write the snapshots straight into the cache buffer using
// the slot stride; otherwise append them after the attn scores (nb_out == 0)
const bool fused = args.nb_out > 0;
const device float * state_out = fused ? (device float *)dst_fuse : (device float *)dst + attn_size;
const uint slot_stride = fused ? (uint)args.nb_out : state_size_per_snap;
for (short t = 0; t < args.ne22; t++) {
float s_k = 0.0f;
@@ -116,7 +123,7 @@ kernel void kernel_gated_delta_net_impl(
if (K > 1) {
const int target_slot = (int)args.ne22 - 1 - (int)t;
if (target_slot >= 0 && target_slot < (int)K) {
device float * dst_state = (device float *) (dst) + attn_size + (uint)target_slot * state_size_per_snap + state_out_base;
device float * dst_state = (device float *)state_out + (uint)target_slot * slot_stride + state_out_base;
FOR_UNROLL (short j = 0; j < NSG; j++) {
const short is = tx*NSG + j;
dst_state[is] = ls[j];
@@ -126,7 +133,7 @@ kernel void kernel_gated_delta_net_impl(
}
if (K == 1) {
device float * dst_state = (device float *) (dst) + attn_size + state_out_base;
device float * dst_state = (device float *)state_out + state_out_base;
FOR_UNROLL (short j = 0; j < NSG; j++) {
const short is = tx*NSG + j;
dst_state[is] = ls[j];
@@ -158,6 +165,7 @@ kernel void kernel_gated_delta_net_impl(
device const char * b,
device const char * s,
device char * dst,
device char * dst_fuse,
uint3 tgpig[[threadgroup_position_in_grid]],
uint3 tpitg[[thread_position_in_threadgroup]],
uint3 ntg[[threads_per_threadgroup]]) {
@@ -230,7 +238,13 @@ kernel void kernel_gated_delta_net_impl(
dst_attn += args.ne21*S_v;
}
device float * dst_state = (device float *) (dst) + args.ne23*args.ne22*args.ne21*S_v + (i23*args.ne21 + i21)*S_v*S_v + i20;
// when fused with the cache cpy, write the snapshots straight into the cache buffer using
// the slot stride; otherwise append them after the attn scores (nb_out == 0)
const bool fused = args.nb_out > 0;
const device float * state_out = fused ? (device float *)dst_fuse : (device float *)dst + args.ne23*args.ne22*args.ne21*S_v;
const uint slot_stride = fused ? (uint)args.nb_out : S_v*S_v;
device float * dst_state = (device float *)state_out + (i23*args.ne21 + i21)*slot_stride + i20;
device T * dstt_state = (device T *) (dst_state);
FOR_UNROLL (short j = 0; j < NSG; j++) {
+64 -35
View File
@@ -496,6 +496,13 @@ kernel void kernel_mul_mm_id(
+ args.nb11*i11
+ args.nb10*iy);
// skip the upper half of the token tile when the expert did not fill it
constexpr short NR1H = NR1/2;
const bool has_hi = nr1 > NR1H;
const short lb1 = (short) tiitg/NL1; // 0 .. NR1-1, this thread's row of the B tile
#ifndef GGML_METAL_HAS_TENSOR
S0_8x8 ma[4];
S1_8x8 mb[2];
@@ -505,15 +512,22 @@ kernel void kernel_mul_mm_id(
for (short i = 0; i < 8; i++){
mc[i] = make_filled_simdgroup_matrix<float, 8>(0.f);
}
// simdgroups 2,3 own rows NR1H..NR1-1
const bool sg_active = has_hi || sgitg < 2;
#else
auto tA = tensor<threadgroup S0, dextents<int32_t, 2>, tensor_inline>(sa, dextents<int32_t, 2>(NK, NR0));
auto tB = tensor<threadgroup S1, dextents<int32_t, 2>, tensor_inline>(sb, dextents<int32_t, 2>(NR1, NK ));
auto tA = tensor<threadgroup S0, dextents<int32_t, 2>, tensor_inline>(sa, dextents<int32_t, 2>(NK, NR0));
// sb is [NR1][NK] row-major
auto tB0 = tensor<threadgroup S1, dextents<int32_t, 2>, tensor_inline>(sb, dextents<int32_t, 2>(NK, NR1H));
auto tB1 = tensor<threadgroup S1, dextents<int32_t, 2>, tensor_inline>(sb + NR1H*NK, dextents<int32_t, 2>(NK, NR1H));
mpp::tensor_ops::matmul2d<
mpp::tensor_ops::matmul2d_descriptor(NR1, NR0, NK, false, true, false, mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate),
mpp::tensor_ops::matmul2d_descriptor(NR1H, NR0, NK, false, true, false, mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate),
execution_simdgroups<4>> mm;
auto cT = mm.get_destination_cooperative_tensor<decltype(tA), decltype(tB), float>();
auto cT0 = mm.get_destination_cooperative_tensor<decltype(tA), decltype(tB0), float>();
auto cT1 = mm.get_destination_cooperative_tensor<decltype(tA), decltype(tB1), float>();
#endif
for (int loop_k = 0; loop_k < args.ne00; loop_k += NK) {
@@ -656,37 +670,45 @@ kernel void kernel_mul_mm_id(
threadgroup_barrier(mem_flags::mem_threadgroup);
#ifndef GGML_METAL_HAS_TENSOR
// load matrices from threadgroup memory and conduct outer products
threadgroup const S0 * lsma = (sa + 4*64*(sgitg%2));
threadgroup const S1 * lsmb = (sb + 2*64*(sgitg/2));
if (sg_active) {
// load matrices from threadgroup memory and conduct outer products
threadgroup const S0 * lsma = (sa + 4*64*(sgitg%2));
threadgroup const S1 * lsmb = (sb + 2*64*(sgitg/2));
FOR_UNROLL (short ik = 0; ik < NK/8; ik++) {
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short ik = 0; ik < NK/8; ik++) {
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 4; i++) {
simdgroup_load(ma[i], lsma + 64*i, 8, 0, false);
FOR_UNROLL (short i = 0; i < 4; i++) {
simdgroup_load(ma[i], lsma + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 2; i++) {
simdgroup_load(mb[i], lsmb + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 8; i++){
simdgroup_multiply_accumulate(mc[i], mb[i/4], ma[i%4], mc[i]);
}
lsma += 8*64;
lsmb += 4*64;
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 2; i++) {
simdgroup_load(mb[i], lsmb + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 8; i++){
simdgroup_multiply_accumulate(mc[i], mb[i/4], ma[i%4], mc[i]);
}
lsma += 8*64;
lsmb += 4*64;
}
#else
auto sA = tA.slice(0, 0);
auto sB = tB.slice(0, 0);
auto sA = tA.slice(0, 0);
auto sB0 = tB0.slice(0, 0);
mm.run(sB, sA, cT);
mm.run(sB0, sA, cT0);
if (has_hi) {
auto sB1 = tB1.slice(0, 0);
mm.run(sB1, sA, cT1);
}
#endif
}
@@ -694,13 +716,20 @@ kernel void kernel_mul_mm_id(
threadgroup_barrier(mem_flags::mem_threadgroup);
#ifdef GGML_METAL_HAS_TENSOR
auto tC = tensor<threadgroup float, dextents<int32_t, 2>, tensor_inline>(sc, dextents<int32_t, 2>(NR0, NR1));
cT.store(tC);
#else
threadgroup float * temp_str = ((threadgroup float *) shmem) + 32*(sgitg&1) + (16*(sgitg >> 1))*NR0;
auto tC0 = tensor<threadgroup float, dextents<int32_t, 2>, tensor_inline>(sc, dextents<int32_t, 2>(NR0, NR1H));
cT0.store(tC0);
for (short i = 0; i < 8; i++) {
simdgroup_store(mc[i], temp_str + 8*(i%4) + 8*NR0*(i/4), NR0, 0, false);
if (has_hi) {
auto tC1 = tensor<threadgroup float, dextents<int32_t, 2>, tensor_inline>(sc + NR1H*NR0, dextents<int32_t, 2>(NR0, NR1H));
cT1.store(tC1);
}
#else
if (sg_active) {
threadgroup float * temp_str = ((threadgroup float *) shmem) + 32*(sgitg&1) + (16*(sgitg >> 1))*NR0;
for (short i = 0; i < 8; i++) {
simdgroup_store(mc[i], temp_str + 8*(i%4) + 8*NR0*(i/4), NR0, 0, false);
}
}
#endif
+226 -82
View File
@@ -1889,8 +1889,19 @@ void kernel_mul_mv_iq2_xxs_f32_impl(
const uint i12 = im%FC_mul_mv_ne12;
const uint i13 = im/FC_mul_mv_ne12;
const uint64_t offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
const int nb32 = nb * (QK_K / 32);
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
const uint64_t offset0 = (first_row + row0)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
device const block_iq2_xxs * x = (device const block_iq2_xxs *) (src0 + offset0);
device const float * y = (device const float *) (src1 + offset1);
@@ -1898,8 +1909,6 @@ void kernel_mul_mv_iq2_xxs_f32_impl(
float yl[32];
float sumf[nr0]={0.f};
const int nb32 = nb * (QK_K / 32);
threadgroup uint64_t * svalues = (threadgroup uint64_t *)(shmem);
threadgroup uint8_t * ssigns = (threadgroup uint8_t *)(svalues + 256);
{
@@ -1912,11 +1921,9 @@ void kernel_mul_mv_iq2_xxs_f32_impl(
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const int ix = tiisg;
device const float * y4 = y + 32 * ix;
for (int ib32 = ix; ib32 < nb32; ib32 += 32) {
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
for (short i = 0; i < 32; ++i) {
yl[i] = y4[i];
}
@@ -1928,7 +1935,7 @@ void kernel_mul_mv_iq2_xxs_f32_impl(
device const uint16_t * q2 = xr->qs + 4 * ib;
device const half * dh = &xr->d;
for (short row = 0; row < nr0; row++) {
for (short row = row0; row < row1; row++) {
const float db = dh[0];
device const uint8_t * aux8 = (device const uint8_t *)q2;
const uint32_t aux32 = q2[2] | (q2[3] << 16);
@@ -1948,7 +1955,7 @@ void kernel_mul_mv_iq2_xxs_f32_impl(
q2 += args.nb01/2;
}
y4 += 32 * 32;
y4 += 32 * ntx;
}
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
@@ -1961,6 +1968,23 @@ void kernel_mul_mv_iq2_xxs_f32_impl(
}
}
template<typename args_t>
void kernel_mul_mv_iq2_xxs_f32_disp(
args_t args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem,
uint3 tgpig,
ushort tiisg,
ushort sgitg) {
if (FC_mul_mv_split) {
kernel_mul_mv_iq2_xxs_f32_impl<N_R0_IQ2_XXS_SPLIT, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_iq2_xxs_f32_impl<N_R0_IQ2_XXS, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
}
[[host_name("kernel_mul_mv_iq2_xxs_f32")]]
kernel void kernel_mul_mv_iq2_xxs_f32(
constant ggml_metal_kargs_mul_mv & args,
@@ -1971,7 +1995,7 @@ kernel void kernel_mul_mv_iq2_xxs_f32(
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_iq2_xxs_f32_impl<N_R0_IQ2_XXS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
kernel_mul_mv_iq2_xxs_f32_disp<constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
template<int nr0, typename args_t>
@@ -1997,8 +2021,19 @@ void kernel_mul_mv_iq2_xs_f32_impl(
const uint i12 = im%FC_mul_mv_ne12;
const uint i13 = im/FC_mul_mv_ne12;
const uint64_t offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
const int nb32 = nb * (QK_K / 32);
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
const uint64_t offset0 = (first_row + row0)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
device const block_iq2_xs * x = (device const block_iq2_xs *) (src0 + offset0);
device const float * y = (device const float *) (src1 + offset1);
@@ -2006,8 +2041,6 @@ void kernel_mul_mv_iq2_xs_f32_impl(
float yl[32];
float sumf[nr0]={0.f};
const int nb32 = nb * (QK_K / 32);
threadgroup uint64_t * svalues = (threadgroup uint64_t *)(shmem);
threadgroup uint8_t * ssigns = (threadgroup uint8_t *)(svalues + 512);
{
@@ -2020,11 +2053,9 @@ void kernel_mul_mv_iq2_xs_f32_impl(
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const int ix = tiisg;
device const float * y4 = y + 32 * ix;
for (int ib32 = ix; ib32 < nb32; ib32 += 32) {
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
for (short i = 0; i < 32; ++i) {
yl[i] = y4[i];
}
@@ -2037,7 +2068,7 @@ void kernel_mul_mv_iq2_xs_f32_impl(
device const uint8_t * sc = xr->scales + ib;
device const half * dh = &xr->d;
for (short row = 0; row < nr0; row++) {
for (short row = row0; row < row1; row++) {
const float db = dh[0];
const uint8_t ls1 = sc[0] & 0xf;
const uint8_t ls2 = sc[0] >> 4;
@@ -2066,7 +2097,7 @@ void kernel_mul_mv_iq2_xs_f32_impl(
sc += args.nb01;
}
y4 += 32 * 32;
y4 += 32 * ntx;
}
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
@@ -2079,6 +2110,23 @@ void kernel_mul_mv_iq2_xs_f32_impl(
}
}
template<typename args_t>
void kernel_mul_mv_iq2_xs_f32_disp(
args_t args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem,
uint3 tgpig,
ushort tiisg,
ushort sgitg) {
if (FC_mul_mv_split) {
kernel_mul_mv_iq2_xs_f32_impl<N_R0_IQ2_XS_SPLIT, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_iq2_xs_f32_impl<N_R0_IQ2_XS, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
}
[[host_name("kernel_mul_mv_iq2_xs_f32")]]
kernel void kernel_mul_mv_iq2_xs_f32(
constant ggml_metal_kargs_mul_mv & args,
@@ -2090,7 +2138,7 @@ kernel void kernel_mul_mv_iq2_xs_f32(
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_iq2_xs_f32_impl<N_R0_IQ2_XS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
kernel_mul_mv_iq2_xs_f32_disp<constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
// FC_mul_mv_split: for nb32 < 32 (nb32 divides 32), 32/nb32 threads share each chunk and each takes a slice of the rows
@@ -2117,8 +2165,19 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
const uint i12 = im%FC_mul_mv_ne12;
const uint i13 = im/FC_mul_mv_ne12;
const uint64_t offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
const int nb32 = nb * (QK_K / 32);
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
const uint64_t offset0 = (first_row + row0)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
device const block_iq3_xxs * x = (device const block_iq3_xxs *) (src0 + offset0);
device const float * y = (device const float *) (src1 + offset1);
@@ -2126,8 +2185,6 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
float yl[32];
float sumf[nr0]={0.f};
const int nb32 = nb * (QK_K / 32);
threadgroup uint32_t * svalues = (threadgroup uint32_t *)(shmem);
threadgroup uint8_t * ssigns = (threadgroup uint8_t *)(svalues + 256);
{
@@ -2140,15 +2197,6 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
device const float * y4 = y + 32 * ix;
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
@@ -2160,9 +2208,9 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
const int ib = ib32 % (QK_K / 32);
device const block_iq3_xxs * xr = x + ibl;
device const uint8_t * q3 = xr->qs + 8 * ib + (uint64_t) row0*args.nb01;
device const uint16_t * gas = (device const uint16_t *)(xr->qs + QK_K/4) + 2 * ib + (uint64_t) row0*args.nb01/2;
device const half * dh = &xr->d + (uint64_t) row0*args.nb01/2;
device const uint8_t * q3 = xr->qs + 8 * ib;
device const uint16_t * gas = (device const uint16_t *)(xr->qs + QK_K/4) + 2 * ib;
device const half * dh = &xr->d;
for (short row = row0; row < row1; row++) {
const float db = dh[0];
@@ -2253,8 +2301,19 @@ void kernel_mul_mv_iq3_s_f32_impl(
const uint i12 = im%FC_mul_mv_ne12;
const uint i13 = im/FC_mul_mv_ne12;
const uint64_t offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
const int nb32 = nb * (QK_K / 32);
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
const uint64_t offset0 = (first_row + row0)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
device const block_iq3_s * x = (device const block_iq3_s *) (src0 + offset0);
device const float * y = (device const float *) (src1 + offset1);
@@ -2262,8 +2321,6 @@ void kernel_mul_mv_iq3_s_f32_impl(
float yl[32];
float sumf[nr0]={0.f};
const int nb32 = nb * (QK_K / 32);
threadgroup uint32_t * svalues = (threadgroup uint32_t *) shmem;
{
int nval = 8;
@@ -2272,11 +2329,9 @@ void kernel_mul_mv_iq3_s_f32_impl(
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const int ix = tiisg;
device const float * y4 = y + 32 * ix;
for (int ib32 = ix; ib32 < nb32; ib32 += 32) {
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
for (short i = 0; i < 32; ++i) {
yl[i] = y4[i];
}
@@ -2291,7 +2346,7 @@ void kernel_mul_mv_iq3_s_f32_impl(
device const uint8_t * signs = xr->signs + 4 * ib;
device const half * dh = &xr->d;
for (short row = 0; row < nr0; row++) {
for (short row = row0; row < row1; row++) {
const float db = dh[0];
const float d = db * (1 + 2*((sc[0] >> 4*(ib%2)) & 0xf));
@@ -2315,7 +2370,7 @@ void kernel_mul_mv_iq3_s_f32_impl(
signs += args.nb01;
}
y4 += 32 * 32;
y4 += 32 * ntx;
}
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
@@ -2328,6 +2383,23 @@ void kernel_mul_mv_iq3_s_f32_impl(
}
}
template<typename args_t>
void kernel_mul_mv_iq3_s_f32_disp(
args_t args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem,
uint3 tgpig,
ushort tiisg,
ushort sgitg) {
if (FC_mul_mv_split) {
kernel_mul_mv_iq3_s_f32_impl<N_R0_IQ3_S_SPLIT, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_iq3_s_f32_impl<N_R0_IQ3_S, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
}
[[host_name("kernel_mul_mv_iq3_s_f32")]]
kernel void kernel_mul_mv_iq3_s_f32(
constant ggml_metal_kargs_mul_mv & args,
@@ -2339,7 +2411,7 @@ kernel void kernel_mul_mv_iq3_s_f32(
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_iq3_s_f32_impl<N_R0_IQ3_S, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
kernel_mul_mv_iq3_s_f32_disp<constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
template<int nr0, typename args_t>
@@ -2365,8 +2437,19 @@ void kernel_mul_mv_iq2_s_f32_impl(
const uint i12 = im%FC_mul_mv_ne12;
const uint i13 = im/FC_mul_mv_ne12;
const uint64_t offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
const int nb32 = nb * (QK_K / 32);
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
const uint64_t offset0 = (first_row + row0)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
device const block_iq2_s * x = (device const block_iq2_s *) (src0 + offset0);
device const float * y = (device const float *) (src1 + offset1);
@@ -2374,8 +2457,6 @@ void kernel_mul_mv_iq2_s_f32_impl(
float yl[32];
float sumf[nr0]={0.f};
const int nb32 = nb * (QK_K / 32);
//threadgroup uint64_t * svalues = (threadgroup uint64_t *) shmem;
//{
// int nval = 32;
@@ -2384,11 +2465,9 @@ void kernel_mul_mv_iq2_s_f32_impl(
// threadgroup_barrier(mem_flags::mem_threadgroup);
//}
const short ix = tiisg;
device const float * y4 = y + 32 * ix;
for (int ib32 = ix; ib32 < nb32; ib32 += 32) {
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
for (short i = 0; i < 32; ++i) {
yl[i] = y4[i];
}
@@ -2403,7 +2482,7 @@ void kernel_mul_mv_iq2_s_f32_impl(
device const uint8_t * signs = qs + QK_K/8;
device const half * dh = &xr->d;
for (short row = 0; row < nr0; row++) {
for (short row = row0; row < row1; row++) {
const float db = dh[0];
const float d1 = db * (0.5f + (sc[0] & 0xf));
const float d2 = db * (0.5f + (sc[0] >> 4));
@@ -2428,7 +2507,7 @@ void kernel_mul_mv_iq2_s_f32_impl(
signs += args.nb01;
}
y4 += 32 * 32;
y4 += 32 * ntx;
}
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
@@ -2441,6 +2520,23 @@ void kernel_mul_mv_iq2_s_f32_impl(
}
}
template<typename args_t>
void kernel_mul_mv_iq2_s_f32_disp(
args_t args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem,
uint3 tgpig,
ushort tiisg,
ushort sgitg) {
if (FC_mul_mv_split) {
kernel_mul_mv_iq2_s_f32_impl<N_R0_IQ2_S_SPLIT, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_iq2_s_f32_impl<N_R0_IQ2_S, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
}
[[host_name("kernel_mul_mv_iq2_s_f32")]]
kernel void kernel_mul_mv_iq2_s_f32(
constant ggml_metal_kargs_mul_mv & args,
@@ -2452,7 +2548,7 @@ kernel void kernel_mul_mv_iq2_s_f32(
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_iq2_s_f32_impl<N_R0_IQ2_S, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
kernel_mul_mv_iq2_s_f32_disp<constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
template<int nr0, typename args_t>
@@ -2478,8 +2574,19 @@ void kernel_mul_mv_iq1_s_f32_impl(
const uint i12 = im%FC_mul_mv_ne12;
const uint i13 = im/FC_mul_mv_ne12;
const uint64_t offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
const int nb32 = nb * (QK_K / 32);
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
const uint64_t offset0 = (first_row + row0)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
device const block_iq1_s * x = (device const block_iq1_s *) (src0 + offset0);
device const float * y = (device const float *) (src1 + offset1);
@@ -2487,13 +2594,9 @@ void kernel_mul_mv_iq1_s_f32_impl(
float yl[32];
float sumf[nr0]={0.f};
const int nb32 = nb * (QK_K / 32);
const short ix = tiisg;
device const float * y4 = y + 32 * ix;
for (int ib32 = ix; ib32 < nb32; ib32 += 32) {
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
float sumy = 0;
for (short i = 0; i < 32; ++i) {
yl[i] = y4[i];
@@ -2508,7 +2611,7 @@ void kernel_mul_mv_iq1_s_f32_impl(
device const uint16_t * qh = xr->qh + ib;
device const half * dh = &xr->d;
for (short row = 0; row < nr0; row++) {
for (short row = row0; row < row1; row++) {
constant uint8_t * grid1 = (constant uint8_t *)(iq1s_grid_gpu + (qs[0] | ((qh[0] << 8) & 0x700)));
constant uint8_t * grid2 = (constant uint8_t *)(iq1s_grid_gpu + (qs[1] | ((qh[0] << 5) & 0x700)));
constant uint8_t * grid3 = (constant uint8_t *)(iq1s_grid_gpu + (qs[2] | ((qh[0] << 2) & 0x700)));
@@ -2528,7 +2631,7 @@ void kernel_mul_mv_iq1_s_f32_impl(
qh += args.nb01/2;
}
y4 += 32 * 32;
y4 += 32 * ntx;
}
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
@@ -2541,6 +2644,23 @@ void kernel_mul_mv_iq1_s_f32_impl(
}
}
template<typename args_t>
void kernel_mul_mv_iq1_s_f32_disp(
args_t args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem,
uint3 tgpig,
ushort tiisg,
ushort sgitg) {
if (FC_mul_mv_split) {
kernel_mul_mv_iq1_s_f32_impl<N_R0_IQ1_S_SPLIT, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_iq1_s_f32_impl<N_R0_IQ1_S, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
}
[[host_name("kernel_mul_mv_iq1_s_f32")]]
kernel void kernel_mul_mv_iq1_s_f32(
constant ggml_metal_kargs_mul_mv & args,
@@ -2551,7 +2671,7 @@ kernel void kernel_mul_mv_iq1_s_f32(
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_iq1_s_f32_impl<N_R0_IQ1_S, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg);
kernel_mul_mv_iq1_s_f32_disp<constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg);
}
template<int nr0, typename args_t>
@@ -2577,8 +2697,19 @@ void kernel_mul_mv_iq1_m_f32_impl(
const uint i12 = im%FC_mul_mv_ne12;
const uint i13 = im/FC_mul_mv_ne12;
const uint64_t offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
const int nb32 = nb * (QK_K / 32);
const short ntx = FC_mul_mv_split ? nb32 : 32;
const short nrep = 32 / ntx;
const short ix = tiisg % ntx;
const short irep = tiisg / ntx;
const short row0 = (nr0 * irep ) / nrep;
const short row1 = (nr0 * (irep + 1)) / nrep;
const uint64_t offset0 = (first_row + row0)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13;
device const block_iq1_m * x = (device const block_iq1_m *) (src0 + offset0);
device const float * y = (device const float *) (src1 + offset1);
@@ -2586,15 +2717,11 @@ void kernel_mul_mv_iq1_m_f32_impl(
float yl[32];
float sumf[nr0]={0.f};
const int nb32 = nb * (QK_K / 32);
const short ix = tiisg;
device const float * y4 = y + 32 * ix;
iq1m_scale_t scale;
for (int ib32 = ix; ib32 < nb32; ib32 += 32) {
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
float4 sumy = {0.f};
for (short i = 0; i < 8; ++i) {
yl[i+ 0] = y4[i+ 0]; sumy[0] += yl[i+ 0];
@@ -2611,7 +2738,7 @@ void kernel_mul_mv_iq1_m_f32_impl(
device const uint8_t * qh = xr->qh + 2 * ib;
device const uint16_t * sc = (device const uint16_t *)xr->scales;
for (short row = 0; row < nr0; row++) {
for (short row = row0; row < row1; row++) {
scale.u16 = (sc[0] >> 12) | ((sc[1] >> 8) & 0x00f0) | ((sc[2] >> 4) & 0x0f00) | (sc[3] & 0xf000);
constant uint8_t * grid1 = (constant uint8_t *)(iq1s_grid_gpu + (qs[0] | ((qh[0] << 8) & 0x700)));
@@ -2637,7 +2764,7 @@ void kernel_mul_mv_iq1_m_f32_impl(
qh += args.nb01;
}
y4 += 32 * 32;
y4 += 32 * ntx;
}
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
@@ -2650,6 +2777,23 @@ void kernel_mul_mv_iq1_m_f32_impl(
}
}
template<typename args_t>
void kernel_mul_mv_iq1_m_f32_disp(
args_t args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem,
uint3 tgpig,
ushort tiisg,
ushort sgitg) {
if (FC_mul_mv_split) {
kernel_mul_mv_iq1_m_f32_impl<N_R0_IQ1_M_SPLIT, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_iq1_m_f32_impl<N_R0_IQ1_M, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
}
[[host_name("kernel_mul_mv_iq1_m_f32")]]
kernel void kernel_mul_mv_iq1_m_f32(
constant ggml_metal_kargs_mul_mv & args,
@@ -2660,7 +2804,7 @@ kernel void kernel_mul_mv_iq1_m_f32(
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_iq1_m_f32_impl<N_R0_IQ1_M, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg);
kernel_mul_mv_iq1_m_f32_disp<constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg);
}
template<int NR0, typename args_t>
@@ -3239,13 +3383,13 @@ template [[host_name("kernel_mul_mv_id_q3_K_f32")]] kernel kernel_mul_mv_id_t
template [[host_name("kernel_mul_mv_id_q4_K_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_q4_K_f32_impl <N_R0_Q4_K>>>;
template [[host_name("kernel_mul_mv_id_q5_K_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_q5_K_f32_impl <N_R0_Q5_K>>>;
template [[host_name("kernel_mul_mv_id_q6_K_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_q6_K_f32_impl <N_R0_Q6_K>>>;
template [[host_name("kernel_mul_mv_id_iq1_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_s_f32_impl <N_R0_IQ1_S>>>;
template [[host_name("kernel_mul_mv_id_iq1_m_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_m_f32_impl <N_R0_IQ1_M>>>;
template [[host_name("kernel_mul_mv_id_iq2_xxs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_xxs_f32_impl<N_R0_IQ2_XXS>>>;
template [[host_name("kernel_mul_mv_id_iq2_xs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_xs_f32_impl <N_R0_IQ2_XS>>>;
template [[host_name("kernel_mul_mv_id_iq1_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_s_f32_disp<ggml_metal_kargs_mul_mv>>>;
template [[host_name("kernel_mul_mv_id_iq1_m_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_m_f32_disp<ggml_metal_kargs_mul_mv>>>;
template [[host_name("kernel_mul_mv_id_iq2_xxs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_xxs_f32_disp<ggml_metal_kargs_mul_mv>>>;
template [[host_name("kernel_mul_mv_id_iq2_xs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_xs_f32_disp<ggml_metal_kargs_mul_mv>>>;
template [[host_name("kernel_mul_mv_id_iq3_xxs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq3_xxs_f32_disp<ggml_metal_kargs_mul_mv>>>;
template [[host_name("kernel_mul_mv_id_iq3_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq3_s_f32_impl <N_R0_IQ3_S>>>;
template [[host_name("kernel_mul_mv_id_iq2_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_s_f32_impl <N_R0_IQ2_S>>>;
template [[host_name("kernel_mul_mv_id_iq3_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq3_s_f32_disp<ggml_metal_kargs_mul_mv>>>;
template [[host_name("kernel_mul_mv_id_iq2_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_s_f32_disp<ggml_metal_kargs_mul_mv>>>;
template [[host_name("kernel_mul_mv_id_iq4_nl_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq4_nl_f32_impl <N_R0_IQ4_NL>>>;
template [[host_name("kernel_mul_mv_id_iq4_xs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq4_xs_f32_impl <N_R0_IQ4_XS>>>;
template [[host_name("kernel_mul_mv_id_tq2_0_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_tq2_0_f32_impl <N_R0_TQ2_0>>>;
+1
View File
@@ -185,6 +185,7 @@ set(GGML_OPENCL_KERNELS
gemv_noshuffle_q4_k_f32_o4
gemv_noshuffle_q4_k_f32_tiled
gemm_noshuffle_q4_k_f32
gemv_noshuffle_q4_k_f32_32b_trans
gemv_noshuffle_q6_k_f32
gemv_noshuffle_q6_k_f32_o4
gemv_noshuffle_q6_k_f32_tiled
+467 -138
View File
@@ -203,39 +203,67 @@ static ggml_cl_version get_opencl_platform_version(cl_platform_id platform) {
return parse_cl_version(param_value);
}
// Returns the DEVICE's OpenCL version. On an error returns ggml_cl_version with all zeroes.
static ggml_cl_version get_opencl_device_version(cl_device_id device) {
size_t param_size;
if (clGetDeviceInfo(device, CL_DEVICE_VERSION, 0, nullptr, &param_size) != CL_SUCCESS || !param_size) {
return {};
}
std::unique_ptr<char[]> param_storage(new char[param_size]);
if (clGetDeviceInfo(device, CL_DEVICE_VERSION, param_size, param_storage.get(), nullptr) != CL_SUCCESS) {
return {};
}
auto param_value = std::string_view(param_storage.get(), param_size);
const std::string version_prefix = "OpenCL "; // "OpenCL <major>.<minor> <device-specific-info>"
if (param_value.find(version_prefix) != 0) {
return {};
}
param_value.remove_prefix(version_prefix.length());
return parse_cl_version(param_value);
}
// Return a version to use in OpenCL C compilation. On an error returns ggml_cl_version with all zeroes.
static ggml_cl_version get_opencl_c_version(ggml_cl_version platform_version, cl_device_id device) {
size_t param_size;
#if CL_TARGET_OPENCL_VERSION >= 300
if (platform_version.major >= 3) {
CL_CHECK(clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_ALL_VERSIONS, 0, nullptr, &param_size));
if (!param_size) {
return {};
// CL_DEVICE_OPENCL_C_ALL_VERSIONS is an OpenCL 3.0 *device* query, so gating it on the
// *platform* version is not enough: a 3.0 platform can expose 2.0 devices, where the
// query returns CL_INVALID_VALUE and the old CL_CHECK aborted during backend init.
// Gate on the device version, and treat a failure as "fall back to the legacy query"
// rather than fatal -- a device may advertise 3.0 and still refuse the property.
const ggml_cl_version device_version = get_opencl_device_version(device);
if (platform_version.major >= 3 && device_version.major >= 3) {
cl_int err = clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_ALL_VERSIONS, 0, nullptr, &param_size);
if (err == CL_SUCCESS && param_size) {
std::unique_ptr<cl_name_version[]> versions(new cl_name_version[param_size]);
err = clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_ALL_VERSIONS, param_size, versions.get(), nullptr);
if (err == CL_SUCCESS) {
unsigned versions_count = param_size / sizeof(cl_name_version);
cl_version version_max = 0;
for (unsigned i = 0; i < versions_count; i++) {
version_max = std::max<cl_version>(versions[i].version, version_max);
}
return { CL_VERSION_MAJOR(version_max), CL_VERSION_MINOR(version_max) };
}
}
std::unique_ptr<cl_name_version[]> versions(new cl_name_version[param_size]);
CL_CHECK(clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_ALL_VERSIONS, param_size, versions.get(), nullptr));
unsigned versions_count = param_size / sizeof(cl_name_version);
cl_version version_max = 0;
for (unsigned i = 0; i < versions_count; i++) {
version_max = std::max<cl_version>(versions[i].version, version_max);
}
return { CL_VERSION_MAJOR(version_max), CL_VERSION_MINOR(version_max) };
// fall through to CL_DEVICE_OPENCL_C_VERSION below
}
#else
GGML_UNUSED(platform_version);
#endif // CL_TARGET_OPENCL_VERSION >= 300
CL_CHECK(clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_VERSION, 0, nullptr, &param_size));
if (!param_size) {
if (clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_VERSION, 0, nullptr, &param_size) != CL_SUCCESS || !param_size) {
return {};
}
std::unique_ptr<char[]> param_storage(new char[param_size]);
CL_CHECK(clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_VERSION, param_size, param_storage.get(), nullptr));
if (clGetDeviceInfo(device, CL_DEVICE_OPENCL_C_VERSION, param_size, param_storage.get(), nullptr) != CL_SUCCESS) {
return {};
}
auto param_value = std::string_view(param_storage.get(), param_size);
const std::string version_prefix = "OpenCL C "; // Suffix: "XX.YY <platform-specific-info>"
@@ -1115,6 +1143,18 @@ struct ggml_backend_opencl_context {
}
void enqueue_ndrange_kernel(cl_kernel kernel, cl_uint work_dim, size_t *global_work_size, size_t *local_work_size, const ggml_tensor * tensor) {
// From the spec on clEnqueueNDRangeKernel:
// If the device associated with command_queue is an OpenCL 2.1 or newer device,
// and global_work_size is NULL or the value in any passed dimension is zero,
// then the kernel command will trivially succeed after its event dependencies
// are satisfied and subsequently update its completion event.
// So this ensures such cases always return trivially without causing errors in
// case of an older device.
for (cl_uint i = 0; i < work_dim; i++) {
if (global_work_size[i] == 0) {
return;
}
}
#ifdef GGML_OPENCL_PROFILING
cl_event evt;
CL_CHECK(clEnqueueNDRangeKernel(queue, kernel, work_dim, NULL, global_work_size, local_work_size, 0, NULL, &evt));
@@ -1185,6 +1225,8 @@ struct ggml_backend_opencl_context {
cl_kernel kernel_convert_block_q4_k_tiled_ns; // tiled-wide convert (opt-in)
cl_kernel kernel_gemv_noshuffle_q4_k_f32_mc3; // multi-column (N=3) verify GEMV
cl_kernel kernel_gemm_noshuffle_q4_k_f32;
cl_kernel kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8_bin;
cl_kernel kernel_gemv_noshuffle_q4_k_f32_32b_trans;
cl_kernel kernel_gemm_noshuffle_q4_k_q8_1_dp4a = nullptr; // dp4a (int8) dense prefill GEMM
cl_kernel kernel_gemm_noshuffle_q4_k_q8_1_dp4a_wimg = nullptr; // dp4a dense prefill GEMM, weights via texture (X1 opt-in)
cl_kernel kernel_gemm_noshuffle_q5_k_q8_1_dp4a = nullptr; // dp4a (int8) dense q5_K prefill GEMM
@@ -4260,6 +4302,43 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
GGML_LOG_CONT(".");
}
backend_ctx->kernel_gemv_noshuffle_q4_k_f32_32b_trans = nullptr;
backend_ctx->kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8_bin = nullptr;
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E) {
{
std::string opts = std::string("-cl-std=") + opencl_c_std +
" -cl-mad-enable "
" -DSIMDGROUP_WIDTH=" +
std::to_string(backend_ctx->adreno_wave_size);
#ifdef GGML_OPENCL_EMBED_KERNELS
const std::string kernel_src {
#include "gemv_noshuffle_q4_k_f32_32b_trans.cl.h"
};
#else
const std::string kernel_src = read_file("gemv_noshuffle_q4_k_f32_32b_trans.cl");
#endif
cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(), opts);
CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q4_k_f32_32b_trans =
clCreateKernel(prog, "gemv_noshuffle_q4_k_f32_32b_trans", &err), err));
CL_CHECK(clReleaseProgram(prog));
GGML_LOG_CONT(".");
}
if (use_adreno_bin_kernels(backend_ctx)) {
size_t bin_size = 0;
const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("gemm_noshuffle_q4_k_f32_32b_trans_ila_a8", &bin_size);
if (kernel_bin && bin_size > 0) {
cl_program bin_prog =
build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, "", bin_size);
CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8_bin =
clCreateKernel(bin_prog, "kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8", &err), err));
CL_CHECK(clReleaseProgram(bin_prog));
GGML_LOG_CONT(".");
}
}
}
std::string CL_moe_compile_opts = std::string("-cl-std=") + opencl_c_std +
" -cl-mad-enable "
" -cl-fast-relaxed-math";
@@ -7722,6 +7801,7 @@ static void ggml_cl_moe_combine_fused(ggml_backend_t backend, const ggml_tensor
}
inline bool use_q4k_tiled(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor); // defined below (used by the GLU-subgraph fuse check)
inline bool use_q4_k_bin_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor);
inline bool use_adreno_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor); // defined below
static bool ggml_opencl_can_fuse(const ggml_backend_opencl_context * backend_ctx, const struct ggml_cgraph * cgraph, int node_idx, std::initializer_list<enum ggml_op> ops) {
@@ -7776,6 +7856,10 @@ static bool ggml_opencl_can_fuse(const ggml_backend_opencl_context * backend_ctx
if (use_q4k_tiled(backend_ctx, gate->src[0]) || use_q4k_tiled(backend_ctx, up->src[0])) {
return false;
}
// q4_K bin kernel requires 32b transposed layout, not compatible with the fused gemv
if (use_q4_k_bin_kernels(backend_ctx, gate->src[0]) || use_q4_k_bin_kernels(backend_ctx, up->src[0])) {
return false;
}
// that noshuffle layout is only produced at set_tensor time when
// use_adreno_kernels() accepts the weight (ne0 >= 512 && ne1 >= 512).
// Smaller weights stay in the plain q4_K layout, which this kernel would
@@ -8349,7 +8433,7 @@ inline bool enable_adreno_trans_weight_q5_K(const ggml_backend_opencl_context *b
qh_img_width <= backend_ctx->image_max_buffer_size;
}
inline bool use_q4_0_ila_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) {
inline bool use_q4_0_bin_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) {
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
if (!backend_ctx->kernel_gemv_noshuffle_q4_0_f32_32b_trans ||
!backend_ctx->kernel_gemm_noshuffle_q4_0_f32_32b_trans_ila_a8_bin) {
@@ -8442,6 +8526,21 @@ static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_backend_opencl_cont
&& tensor->ne[2] == 1 && tensor->ne[3] == 1;
}
inline bool use_q4_k_bin_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) {
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
if (!backend_ctx->kernel_gemv_noshuffle_q4_k_f32_32b_trans ||
!backend_ctx->kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8_bin) {
return false;
}
return (tensor->ne[0] % 256 == 0) && (tensor->ne[1] % 64 == 0) &&
!use_q4k_tiled(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q4_K(backend_ctx, tensor);
#else
GGML_UNUSED(backend_ctx);
GGML_UNUSED(tensor);
return false;
#endif
}
static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_tensor * op) {
ggml_backend_opencl_device_context * dev_ctx = (ggml_backend_opencl_device_context *)dev->context;
ggml_backend_opencl_context * backend_ctx = dev_ctx->backend_ctx;
@@ -9400,6 +9499,96 @@ static enum ggml_status ggml_backend_opencl_buffer_init_tensor(ggml_backend_buff
return GGML_STATUS_SUCCESS;
}
// Allocate a temporary upload buffer of `nbytes` and populate it with `data`
// from host. On Adreno X1-85 the device-only pool intermittently fails to
// allocate at hundreds of MB once model weights fragment the heap (observed
// on Qwen3.5-9B output.weight Q6_K at 834 MB). Three-step retry:
// 1. CL_MEM_READ_WRITE alloc + clEnqueueWriteBuffer (normal fast path).
// 2. clFinish + retry (drains in-flight allocs that may be holding heap;
// mirrors the proven pattern at the FD-split partial buffer alloc).
// 3. CL_MEM_ALLOC_HOST_PTR + map(WRITE_INVALIDATE) + memcpy + unmap —
// different memory pool (host-pinned); true zero-copy on Adreno per
// QCOM guidance. (CL_MEM_USE_HOST_PTR is NOT zero-copy on Adreno: the
// driver triggers an internal copy because arbitrary host pages aren't
// guaranteed mappable/coherent, AND it draws from the same exhausted
// device pool — so it doesn't solve the problem.)
// Returns the ready-to-read buffer (caller must clReleaseMemObject) or NULL
// if all three strategies fail. The buffer is opaque to the caller — it can
// be passed as a kernel argument like any normal cl_mem.
static cl_mem ggml_cl_create_temp_upload_buffer(
cl_context context, cl_command_queue queue,
size_t nbytes, const void * data,
const char * tensor_name_for_log)
{
cl_int err;
cl_mem buf = clCreateBuffer(context, CL_MEM_READ_WRITE, nbytes, NULL, &err);
if (err != CL_SUCCESS) {
clFinish(queue);
buf = clCreateBuffer(context, CL_MEM_READ_WRITE, nbytes, NULL, &err);
}
if (err == CL_SUCCESS) {
const cl_int werr = clEnqueueWriteBuffer(queue, buf, CL_TRUE, 0, nbytes, data, 0, NULL, NULL);
if (werr == CL_SUCCESS) {
return buf;
}
clReleaseMemObject(buf);
}
buf = clCreateBuffer(context,
CL_MEM_READ_ONLY | CL_MEM_ALLOC_HOST_PTR | CL_MEM_HOST_WRITE_ONLY,
nbytes, NULL, &err);
if (err != CL_SUCCESS) {
return NULL;
}
void * mapped = clEnqueueMapBuffer(queue, buf, CL_TRUE,
CL_MAP_WRITE_INVALIDATE_REGION, 0, nbytes, 0, NULL, NULL, &err);
if (err != CL_SUCCESS) {
clReleaseMemObject(buf);
return NULL;
}
memcpy(mapped, data, nbytes);
const cl_int uerr = clEnqueueUnmapMemObject(queue, buf, mapped, 0, NULL, NULL);
if (uerr != CL_SUCCESS) {
clReleaseMemObject(buf);
return NULL;
}
if (tensor_name_for_log) {
GGML_LOG_INFO("ggml_opencl: %s (%.1f MiB) — device alloc failed, using CL_MEM_ALLOC_HOST_PTR fallback\n",
tensor_name_for_log, nbytes / 1024.0 / 1024.0);
}
return buf;
}
// Allocate a temporary download buffer of `nbytes`. The caller runs a kernel
// that writes into it, then reads it back to host via clEnqueueReadBuffer (or
// equivalent). Mirrors ggml_cl_create_temp_upload_buffer; the host-pinned
// fallback flags are flipped (CL_MEM_WRITE_ONLY | HOST_READ_ONLY) and the
// helper doesn't populate the buffer.
static cl_mem ggml_cl_create_temp_download_buffer(
cl_context context, cl_command_queue queue,
size_t nbytes, const char * tensor_name_for_log)
{
cl_int err;
cl_mem buf = clCreateBuffer(context, CL_MEM_READ_WRITE, nbytes, NULL, &err);
if (err != CL_SUCCESS) {
clFinish(queue);
buf = clCreateBuffer(context, CL_MEM_READ_WRITE, nbytes, NULL, &err);
}
if (err == CL_SUCCESS) {
return buf;
}
buf = clCreateBuffer(context,
CL_MEM_WRITE_ONLY | CL_MEM_ALLOC_HOST_PTR | CL_MEM_HOST_READ_ONLY,
nbytes, NULL, &err);
if (err != CL_SUCCESS) {
return NULL;
}
if (tensor_name_for_log) {
GGML_LOG_INFO("ggml_opencl: %s download (%.1f MiB) — device alloc failed, using CL_MEM_ALLOC_HOST_PTR fallback\n",
tensor_name_for_log, nbytes / 1024.0 / 1024.0);
}
return buf;
}
static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, size_t offset, size_t size) {
ggml_backend_opencl_device_context * dev_ctx = (ggml_backend_opencl_device_context *) buffer->buft->device->context;
ggml_backend_opencl_context * backend_ctx = dev_ctx->backend_ctx;
@@ -9508,12 +9697,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_d + size_q == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "set_tensor: temp upload buffer alloc failed");
// We consider the specified offset arg as always, although For weights
// the offset arg should be 0 (we do not assert this).
@@ -9625,7 +9810,7 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(K % 32 == 0);
if (use_q4_0_ila_kernels(backend_ctx, tensor)) {
if (use_q4_0_bin_kernels(backend_ctx, tensor)) {
cl_int err;
cl_image_format wimg_fmt;
cl_image_desc wimg_desc;
@@ -9671,12 +9856,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_d + size_m + size_q == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "set_tensor: temp upload buffer alloc failed");
cl_buffer_region region;
@@ -9803,12 +9984,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_d + size_qs + size_qh == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "set_tensor: temp upload buffer alloc failed");
cl_buffer_region region;
@@ -9967,12 +10144,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_d + size_m + size_qs + size_qh == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "set_tensor: temp upload buffer alloc failed");
cl_buffer_region region;
@@ -10120,12 +10293,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_e + size_q == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "set_tensor: temp upload buffer alloc failed");
// The original tensor memory is divided into scales and quants, i.e.,
// we first store scales, then quants.
@@ -10231,12 +10400,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_d + size_q == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "set_tensor: temp upload buffer alloc failed");
// The original tensor memory is divided into scales and quants, i.e.,
// we first store scales, then quants.
@@ -10335,12 +10500,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_d + size_q == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "set_tensor: temp upload buffer alloc failed");
cl_buffer_region region;
@@ -10419,12 +10580,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(size_d + size_dm + size_s + size_q == ggml_nbytes(tensor) && "Incorrect tensor size");
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
CL_CHECK(clEnqueueWriteBuffer(
queue, data_device, CL_TRUE, 0,
ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "q4_K set_tensor: temp upload buffer alloc failed");
cl_buffer_region region;
@@ -10576,8 +10733,25 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
GGML_ASSERT(K % 32 == 0);
// Transpose q, d, dm as ushort
transpose_2d_as_16b(backend_ctx, extra->q, extra->q, size_q, K/4, M);
if (use_q4_k_bin_kernels(backend_ctx, tensor)) {
cl_int err;
cl_image_format wimg_fmt;
cl_image_desc wimg_desc;
// transpose quants as 32-bit words (M-first)
GGML_ASSERT(M % 64 == 0);
transpose_2d_as_32b(backend_ctx, extra->q, extra->q, size_q, K/8, M);
wimg_fmt = { CL_R, CL_UNSIGNED_INT32 };
memset(&wimg_desc, 0, sizeof(wimg_desc));
wimg_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
wimg_desc.image_width = (size_t)M * K / 8;
wimg_desc.buffer = extra->q;
CL_CHECK((extra->q_img = clCreateImage(context, CL_MEM_READ_ONLY, &wimg_fmt, &wimg_desc, NULL, &err), err));
} else {
// Transpose q as ushort
transpose_2d_as_16b(backend_ctx, extra->q, extra->q, size_q, K/4, M);
}
transpose_2d_as_16b(backend_ctx, extra->d, extra->d, size_d, K/256, M);
transpose_2d_as_16b(backend_ctx, extra->dm, extra->dm, size_dm, K/256, M);
@@ -10604,9 +10778,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
"Incorrect tensor size");
cl_int err;
cl_mem data_device;
CL_CHECK((data_device = clCreateBuffer(context, CL_MEM_READ_WRITE, ggml_nbytes(tensor), NULL, &err), err));
CL_CHECK(clEnqueueWriteBuffer(queue, data_device, CL_TRUE, 0, ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "q5_K set_tensor: temp upload buffer alloc failed");
cl_buffer_region region;
@@ -10792,9 +10965,8 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
"Incorrect tensor size");
cl_int err;
cl_mem data_device;
CL_CHECK((data_device = clCreateBuffer(context, CL_MEM_READ_WRITE, ggml_nbytes(tensor), NULL, &err), err));
CL_CHECK(clEnqueueWriteBuffer(queue, data_device, CL_TRUE, 0, ggml_nbytes(tensor), data, 0, NULL, NULL));
cl_mem data_device = ggml_cl_create_temp_upload_buffer(context, queue, ggml_nbytes(tensor), data, tensor->name);
GGML_ASSERT(data_device != NULL && "q6_K set_tensor: temp upload buffer alloc failed");
cl_buffer_region region;
@@ -11135,9 +11307,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
cl_int err;
cl_kernel kernel = backend_ctx->kernel_restore_block_q4_0_trans4_ns;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
int ne00 = tensor->ne[0];
int ne01 = tensor->ne[1];
@@ -11180,7 +11351,7 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
buf_trans_d.allocate(backend_ctx->context, size_d);
buf_unpacked.allocate(backend_ctx->context, ggml_nbytes(tensor));
if (use_q4_0_ila_kernels(backend_ctx, tensor)) {
if (use_q4_0_bin_kernels(backend_ctx, tensor)) {
transpose_2d_as_32b(backend_ctx, extra->q, buf_trans_q.buffer, size_q, M, K / 8);
} else {
transpose_2d_as_16b(backend_ctx, extra->q, buf_trans_q.buffer, size_q, M, K / 4);
@@ -11206,10 +11377,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
}
#endif
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q4_0;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra->q));
@@ -11234,10 +11403,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
if (use_adreno_moe_kernels(backend_ctx, tensor)) {
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q4_1_trans4_ns;
int ne00 = tensor->ne[0];
@@ -11309,10 +11476,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
}
#endif
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q4_1;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra->q));
@@ -11340,9 +11505,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
if (use_adreno_moe_kernels(backend_ctx, tensor)) {
cl_int err;
// TODO: use ggml_cl_buffer to manage this temporary buffer
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q5_0_trans4_ns;
@@ -11444,9 +11608,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
if (use_adreno_moe_kernels(backend_ctx, tensor)) {
cl_int err;
// TODO: use ggml_cl_buffer to manage this temporary buffer
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q5_1_trans4_ns;
@@ -11551,10 +11714,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
if (tensor->type == GGML_TYPE_MXFP4) {
ggml_tensor_extra_cl_mxfp4 * extra = (ggml_tensor_extra_cl_mxfp4 *)tensor->extra;
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
if (use_adreno_moe_kernels(backend_ctx, tensor)) {
@@ -11616,10 +11777,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
const ggml_tensor * extra_src = tensor->view_src != nullptr ? tensor->view_src : tensor;
ggml_tensor_extra_cl_q8_0 * extra = (ggml_tensor_extra_cl_q8_0 *)extra_src->extra;
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
if (enable_adreno_trans_weight(backend_ctx, tensor)) {
@@ -11672,10 +11831,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
if (tensor->type == GGML_TYPE_IQ4_NL) {
ggml_tensor_extra_cl_iq4_nl * extra = (ggml_tensor_extra_cl_iq4_nl *)tensor->extra;
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
if (use_adreno_kernels(backend_ctx, tensor)) {
@@ -11744,10 +11901,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
if (tensor->type == GGML_TYPE_Q4_K) {
ggml_tensor_extra_cl_q4_K * extra = (ggml_tensor_extra_cl_q4_K *)tensor->extra;
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_uchar mask_0F = 0x0F;
cl_uchar mask_F0 = 0xF0;
@@ -11802,10 +11957,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
return;
}
if (use_adreno_moe_kernels(backend_ctx, tensor)) {
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q4_k_trans4_ns;
@@ -11855,7 +12008,11 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
buf_trans_s.allocate(backend_ctx->context, size_s);
// Transpose q, d, dm, s back
transpose_2d_as_16b(backend_ctx, extra->q, buf_trans_q.buffer, size_q, M, K/4);
if (use_q4_k_bin_kernels(backend_ctx, tensor)) {
transpose_2d_as_32b(backend_ctx, extra->q, buf_trans_q.buffer, size_q, M, K/8);
} else {
transpose_2d_as_16b(backend_ctx, extra->q, buf_trans_q.buffer, size_q, M, K/4);
}
transpose_2d_as_16b(backend_ctx, extra->d, buf_trans_d.buffer, size_d, M, K/256);
transpose_2d_as_16b(backend_ctx, extra->dm, buf_trans_dm.buffer, size_dm, M, K/256);
transpose_2d_as_8b (backend_ctx, extra->s, buf_trans_s.buffer, size_s, M, K/256*12, true, true);
@@ -11906,20 +12063,16 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
if (tensor->type == GGML_TYPE_Q5_K) {
ggml_tensor_extra_cl_q5_K * extra = (ggml_tensor_extra_cl_q5_K *)tensor->extra;
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_uchar mask_0F = 0x0F;
cl_uchar mask_F0 = 0xF0;
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
if (use_adreno_moe_kernels(backend_ctx, tensor)) {
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q5_k_trans4_ns;
int ne00 = tensor->ne[0];
@@ -12079,10 +12232,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
return;
}
if (use_adreno_moe_kernels(backend_ctx, tensor)) {
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_kernel kernel = backend_ctx->kernel_restore_block_q6_k_trans4_ns;
@@ -12169,10 +12320,8 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
}
#endif // GGML_OPENCL_USE_ADRENO_KERNELS
cl_int err;
cl_mem data_device = clCreateBuffer(context, CL_MEM_READ_WRITE,
ggml_nbytes(tensor), NULL, &err);
CL_CHECK(err);
cl_mem data_device = ggml_cl_create_temp_download_buffer(context, queue, ggml_nbytes(tensor), tensor->name);
GGML_ASSERT(data_device != NULL && "get_tensor: temp download buffer alloc failed");
cl_uchar mask = 0xFF;
cl_ulong n_blk = ggml_nelements(tensor)/ggml_blck_size(tensor->type);
@@ -12300,6 +12449,21 @@ static ggml_backend_buffer_t ggml_backend_opencl_buffer_type_alloc_buffer(ggml_b
cl_int err;
cl_mem mem = clCreateBuffer(backend_ctx->context, CL_MEM_READ_WRITE, size, NULL, &err);
// On Adreno X1-85 the device pool intermittently fails at hundreds of MB
// once the heap fragments (e.g. graph-allocator compute-buffer reserve
// after model load). Four-step retry:
// 1. normal alloc (fast path)
// 2. clFinish + retry (drains in-flight allocs)
// 3. cl_qcom_large_buffer (X2-class driver only, OpenCL 3.0 only)
// 4. ALLOC_HOST_PTR (host-pinned pool) — last-resort fallback. This
// buffer backs compute scratch read/written by every kernel in the
// graph, so kernel accesses fall to host memory and runtime perf
// degrades meaningfully. Better than failing to load, but the user
// should see the warning and consider -ngl reduction.
if (err != CL_SUCCESS) {
clFinish(backend_ctx->queue);
mem = clCreateBuffer(backend_ctx->context, CL_MEM_READ_WRITE, size, NULL, &err);
}
#if GGML_OPENCL_TARGET_VERSION >= 300
// clCreateBufferWithProperties and cl_mem_properties are OpenCL 3.0. Drivers older than
// that do not export the symbol, so a build targeting them fails to link. The large
@@ -12310,9 +12474,20 @@ static ggml_backend_buffer_t ggml_backend_opencl_buffer_type_alloc_buffer(ggml_b
mem = clCreateBufferWithProperties(backend_ctx->context, props, CL_MEM_READ_WRITE, size, NULL, &err);
}
#endif
if (err != CL_SUCCESS) {
mem = clCreateBuffer(backend_ctx->context, CL_MEM_READ_WRITE | CL_MEM_ALLOC_HOST_PTR, size, NULL, &err);
if (err == CL_SUCCESS) {
GGML_LOG_WARN("%s: %.2f MiB allocated via CL_MEM_ALLOC_HOST_PTR fallback — "
"device pool exhausted; runtime perf will be degraded. "
"Consider lowering -ngl or context size.\n",
__func__, size / 1024.0 / 1024.0);
}
}
if (err != CL_SUCCESS) {
GGML_LOG_INFO("%s: failed to allocate %.2f MiB\n", __func__, size / 1024.0 / 1024.0);
GGML_LOG_ERROR("%s: failed to allocate %.2f MiB (err=%d). "
"Consider reducing -ngl, lowering -c / -ub, or using quantized KV cache.\n",
__func__, size / 1024.0 / 1024.0, err);
return nullptr;
}
@@ -13067,6 +13242,7 @@ static void ggml_cl_set_rows(ggml_backend_t backend, const ggml_tensor * src0, c
(size_t)ne03};
size_t local_work_size[] = {(size_t)nth0, (size_t)rows_per_workgroup, 1};
// ne01 == 0 makes global_work_size[0] zero here; enqueue_ndrange_kernel drops the empty range.
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
}
@@ -18639,9 +18815,9 @@ static void ggml_cl_mul_mat_q4_0_f32_adreno(ggml_backend_t backend, const ggml_t
static const bool q40_mc3 = (getenv("GGML_OPENCL_Q40_MC3") != nullptr);
const bool use_q40_mc3 = q40_mc3 && (ne1 >= 2 && ne1 <= 4) && (ne01 < 32768);
const bool use_ila = use_q4_0_ila_kernels(backend_ctx, src0);
const bool use_bin = use_q4_0_bin_kernels(backend_ctx, src0);
if (use_ila) {
if (use_bin) {
if (use_q40_mc3) {
static bool warned = false;
if (!warned) {
@@ -20196,6 +20372,145 @@ static void ggml_cl_mul_mat_q8_0_f32_adreno(ggml_backend_t backend, const ggml_t
#endif
}
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
static void ggml_cl_mul_mat_q4_k_f32_adreno_ila(ggml_backend_t backend, const ggml_tensor * src0,
const ggml_tensor * src1, ggml_tensor * dst) {
GGML_ASSERT(src0);
GGML_ASSERT(src0->extra);
GGML_ASSERT(src1);
GGML_ASSERT(src1->extra);
GGML_ASSERT(dst);
GGML_ASSERT(dst->extra);
ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context;
ggml_tensor_extra_cl * extra1 = (ggml_tensor_extra_cl *)src1->extra;
ggml_tensor_extra_cl * extrad = (ggml_tensor_extra_cl *)dst->extra;
ggml_tensor_extra_cl_q4_K * extra0_q4_k = (ggml_tensor_extra_cl_q4_K *)src0->extra;
cl_ulong offset1 = extra1->offset + src1->view_offs;
cl_ulong offsetd = extrad->offset + dst->view_offs;
const int ne00 = src0->ne[0];
const int ne01 = src0->ne[1];
const int ne1 = dst->ne[1];
GGML_ASSERT(ne00 % ggml_blck_size(src0->type) == 0);
cl_context context = backend_ctx->context;
cl_kernel kernel;
cl_int err;
cl_image_format img_fmt;
cl_image_desc img_desc;
cl_buffer_region region;
int M = ne01;
int N = ne1;
int K = ne00;
if (ne1 == 1) {
cl_mem b_sub_buf = nullptr;
cl_mem b_img = nullptr;
region.origin = offset1;
region.size = (size_t)K * N * sizeof(float);
CL_CHECK((b_sub_buf = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, &region, &err), err));
img_fmt = { CL_RGBA, CL_FLOAT };
memset(&img_desc, 0, sizeof(img_desc));
img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
img_desc.image_width = (size_t)K * N / 4;
img_desc.buffer = b_sub_buf;
CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err));
kernel = backend_ctx->kernel_gemv_noshuffle_q4_k_f32_32b_trans;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q4_k->q_img));
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_k->d));
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q4_k->dm));
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q4_k->s));
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img));
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &extrad->data_device));
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_ulong), &offsetd));
CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne00));
CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne01));
size_t local_work_size[3] = { 64, 8, 1 };
size_t global_work_size[3] = { (size_t)ne01, 8, 1 };
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
CL_CHECK(clReleaseMemObject(b_sub_buf));
CL_CHECK(clReleaseMemObject(b_img));
} else {
const int gemm_tile_n = 64;
int N_pad = CEIL_DIV(N, gemm_tile_n) * gemm_tile_n;
cl_mem b_sub_buf = nullptr;
cl_mem b_padded = nullptr;
cl_mem b_buf = nullptr;
if (N_pad == N) {
region.origin = offset1;
region.size = (size_t)K * N * sizeof(float);
CL_CHECK((b_sub_buf = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, &region, &err), err));
b_buf = b_sub_buf;
} else {
CL_CHECK((b_padded = clCreateBuffer(context, CL_MEM_READ_WRITE, (size_t)K * N_pad * sizeof(float), NULL, &err), err));
const float zero = 0.0f;
CL_CHECK(clEnqueueFillBuffer(backend_ctx->queue, b_padded, &zero, sizeof(zero), 0, (size_t)K * N_pad * sizeof(float), 0, NULL, NULL));
CL_CHECK(clEnqueueCopyBuffer(backend_ctx->queue, extra1->data_device, b_padded, offset1, 0, (size_t)K * N * sizeof(float), 0, NULL, NULL));
b_buf = b_padded;
}
img_fmt = { CL_R, CL_FLOAT };
memset(&img_desc, 0, sizeof(img_desc));
img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
img_desc.image_width = (size_t)K * N_pad;
img_desc.buffer = b_buf;
cl_mem b_img;
CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err));
region.origin = offsetd;
region.size = (size_t)M * N * sizeof(float);
cl_mem d_sub_buf;
CL_CHECK((d_sub_buf = clCreateSubBuffer(extrad->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, &region, &err), err));
img_fmt = { CL_R, CL_FLOAT };
memset(&img_desc, 0, sizeof(img_desc));
img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
img_desc.image_width = (size_t)M * N;
img_desc.buffer = d_sub_buf;
cl_mem d_img;
CL_CHECK((d_img = clCreateImage(context, CL_MEM_WRITE_ONLY, &img_fmt, &img_desc, NULL, &err), err));
kernel = backend_ctx->kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8_bin;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q4_k->q_img));
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q4_k->d));
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q4_k->dm));
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q4_k->s));
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img));
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &d_img));
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_uint), &ne00));
CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_uint), &ne01));
CL_CHECK(clSetKernelArg(kernel, 8, sizeof(int), &N));
size_t local_work_size[3] = { 64, 2, 2 };
size_t m_tiles = (size_t)CEIL_DIV(M, 64);
size_t global_work_size[3] = { 64, m_tiles, (size_t)CEIL_DIV(N_pad, gemm_tile_n) };
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
CL_CHECK(clReleaseMemObject(b_img));
if (b_sub_buf) {
CL_CHECK(clReleaseMemObject(b_sub_buf));
}
if (b_padded) {
CL_CHECK(clReleaseMemObject(b_padded));
}
CL_CHECK(clReleaseMemObject(d_img));
CL_CHECK(clReleaseMemObject(d_sub_buf));
}
}
#endif // GGML_OPENCL_USE_ADRENO_KERNELS
static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
GGML_ASSERT(src0);
@@ -20248,6 +20563,20 @@ static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_t
// unified routes batched Q6_K lm_head to CPU). Per-layer mc3 is byte-identical.
const bool use_mc3 = q4k_mc3 && (ne1 == 3) && (ne01 < 32768);
const bool use_bin = use_q4_k_bin_kernels(backend_ctx, src0);
if (use_bin) {
if (use_mc3) {
static bool warned = false;
if (!warned) {
GGML_LOG_WARN("ggml_opencl: GGML_OPENCL_Q4K_MC3 is bypassed by Q4_K binary kernels\n");
warned = true;
}
}
ggml_cl_mul_mat_q4_k_f32_adreno_ila(backend, src0, src1, dst);
return;
}
if (ne1 == 1 || use_mc3) {
cl_mem q_img = nullptr;
cl_mem b_sub_buf = nullptr;
@@ -0,0 +1,134 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
#define QK_K 256
#define K_SCALE_SIZE 12
#define N_SIMDGROUP 8
#define SIMDGROUP_WIDTH 64
inline void get_scale_min_k4(
int j,
global const uchar * q,
uint stride,
uchar * d,
uchar * m
) {
if (j < 4) {
*d = q[j*stride] & 63;
*m = q[(j+4)*stride] & 63;
} else {
*d = (q[(j+4)*stride] & 0x0F) | ((q[(j-4)*stride] & 0xC0) >> 2);
*m = ((q[(j+4)*stride] >> 4) & 0x0F) | ((q[j*stride] & 0xC0) >> 2);
}
}
static inline float8 q4_k_to_fp32_packed8(ushort2 q4x8, float scale, float minv) {
float8 fp32x8;
fp32x8.s0 = (q4x8.s0 & 0x000F) * scale - minv;
fp32x8.s1 = ((q4x8.s0 & 0x00F0) >> 4) * scale - minv;
fp32x8.s2 = ((q4x8.s0 & 0x0F00) >> 8) * scale - minv;
fp32x8.s3 = ((q4x8.s0 & 0xF000) >> 12) * scale - minv;
fp32x8.s4 = (q4x8.s1 & 0x000F) * scale - minv;
fp32x8.s5 = ((q4x8.s1 & 0x00F0) >> 4) * scale - minv;
fp32x8.s6 = ((q4x8.s1 & 0x0F00) >> 8) * scale - minv;
fp32x8.s7 = ((q4x8.s1 & 0xF000) >> 12) * scale - minv;
return fp32x8;
}
__attribute__((qcom_reqd_sub_group_size("half")))
__kernel void gemv_noshuffle_q4_k_f32_32b_trans(
read_only image1d_buffer_t src0_q,
__global half * src0_d,
__global half * src0_dm,
__global uchar * src0_s,
__read_only image1d_buffer_t src1,
__global float * dst,
ulong offsetd,
int ne00,
int ne01
) {
uint i01 = get_global_id(0);
uint sgid = get_local_id(1);
uint slid = get_sub_group_local_id();
int num_subblocks = ne00 / 32;
__private float sum = 0.0f;
// Loop over sub-blocks of 32 elements, N_SIMDGROUP sub-blocks per iter
for (uint ib = sgid; ib < num_subblocks; ib += N_SIMDGROUP) {
uint sb = ib / 8;
uint j = ib % 8;
// Load d and dmin for this super-block
half d_val = src0_d[sb * ne01 + i01];
half dm_val = src0_dm[sb * ne01 + i01];
// Load sub-block scale and min. s is transposed [nb][12][M]; stride ne01 per code.
global const uchar * sc = src0_s + sb * K_SCALE_SIZE * ne01 + i01;
uchar sv, mn;
get_scale_min_k4(j, sc, ne01, &sv, &mn);
float scale = (float)d_val * (float)sv;
float minv = (float)dm_val * (float)mn;
// Load 4 uints of quants (32 nibbles = 32 elements), column-major stride ne01
uint q_base = ib * ne01 * 4 + i01;
uint4 regQ;
regQ.s0 = read_imageui(src0_q, q_base).x;
regQ.s1 = read_imageui(src0_q, q_base + ne01).x;
regQ.s2 = read_imageui(src0_q, q_base + ne01 * 2).x;
regQ.s3 = read_imageui(src0_q, q_base + ne01 * 3).x;
// Load activations: 32 floats = 8 float4s
uint y_offset = ib * 8;
float4 y_local = (slid < 8) ? read_imagef(src1, (y_offset + slid)) : (float4)0.0f;
float4 y0 = sub_group_broadcast(y_local, 0);
float4 y1 = sub_group_broadcast(y_local, 1);
float4 y2 = sub_group_broadcast(y_local, 2);
float4 y3 = sub_group_broadcast(y_local, 3);
float4 y4 = sub_group_broadcast(y_local, 4);
float4 y5 = sub_group_broadcast(y_local, 5);
float4 y6 = sub_group_broadcast(y_local, 6);
float4 y7 = sub_group_broadcast(y_local, 7);
float8 fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s0), scale, minv);
float4 acc = y0 * fp32x8.lo;
acc += y1 * fp32x8.hi;
fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s1), scale, minv);
acc += y2 * fp32x8.lo;
acc += y3 * fp32x8.hi;
fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s2), scale, minv);
acc += y4 * fp32x8.lo;
acc += y5 * fp32x8.hi;
fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s3), scale, minv);
acc += y6 * fp32x8.lo;
acc += y7 * fp32x8.hi;
sum += ((acc.s0 + acc.s1) + (acc.s2 + acc.s3));
}
// reduction in local memory over N_SIMDGROUP subgroups
__local float reduceLM[SIMDGROUP_WIDTH * (N_SIMDGROUP - 1)];
if (sgid > 0) {
reduceLM[SIMDGROUP_WIDTH * (sgid - 1) + slid] = sum;
}
barrier(CLK_LOCAL_MEM_FENCE);
if (sgid == 0) {
for (uint i = 0; i < N_SIMDGROUP - 1; ++i) {
sum += reduceLM[SIMDGROUP_WIDTH * i + slid];
}
}
// 1 output per thread in subgroup 0
if (sgid == 0) {
dst = dst + (offsetd >> 2);
dst[i01] = sum;
}
}
+4 -2
View File
@@ -36,8 +36,10 @@ if (GGML_RPC_RDMA)
target_compile_definitions(ggml-rpc PRIVATE GGML_RPC_RDMA)
if (APPLE)
# librdma.dylib only exists on macOS 26.2 and later. Link it weakly so a build made
# where it exists still loads where it does not; checked at runtime before use.
target_link_options(ggml-rpc PRIVATE "LINKER:-weak_library,${RDMA_LIB}")
# where it exists still loads where it does not; checked at runtime before use
# but with BUILD_SHARED_LIBS=OFF ggml-rpc is a static archive and never links
# so the librdma symbols used by transport-apple.cpp stay undefined.
target_link_options(ggml-rpc PUBLIC "LINKER:-weak_library,${RDMA_LIB}")
target_compile_definitions(ggml-rpc PRIVATE GGML_RPC_RDMA_APPLE)
target_sources(ggml-rpc PRIVATE transport-apple.cpp)
else()
+6 -4
View File
@@ -6326,7 +6326,7 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons
return false;
}
if (src0_type == GGML_TYPE_TQ2_0) {
if (src0_type == GGML_TYPE_TQ2_0 || src0_type == GGML_TYPE_TQ1_0) {
return false;
}
@@ -6380,7 +6380,7 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons
case GGML_OP_SET_ROWS:
{
if (op->type == GGML_TYPE_TQ2_0) {
if (op->type == GGML_TYPE_TQ2_0 || op->type == GGML_TYPE_TQ1_0) {
return false;
}
auto res = (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16 ||
@@ -6502,12 +6502,14 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons
src1_type == GGML_TYPE_IQ3_S ||
src1_type == GGML_TYPE_IQ1_S ||
src1_type == GGML_TYPE_IQ1_M ||
src1_type == GGML_TYPE_TQ2_0) {
src1_type == GGML_TYPE_TQ2_0 ||
src1_type == GGML_TYPE_TQ1_0) {
return false;
}
}
if (src0_type == GGML_TYPE_TQ2_0 || src1_type == GGML_TYPE_TQ2_0) {
if (src0_type == GGML_TYPE_TQ2_0 || src1_type == GGML_TYPE_TQ2_0 ||
src0_type == GGML_TYPE_TQ1_0 || src1_type == GGML_TYPE_TQ1_0) {
return false;
}
+6
View File
@@ -39,6 +39,12 @@ ggml_add_backend_library(ggml-webgpu
add_dependencies(ggml-webgpu generate_shaders)
# Dawn needs C++20 (https://dawn.googlesource.com/dawn/+/refs/heads/main/docs/quickstart-cmake.md#prerequisites)
target_compile_features(ggml-webgpu PRIVATE cxx_std_20)
# Disable C++20 module scanning since emscan-deps fails to find webgpu_cpp.h
set_target_properties(ggml-webgpu PROPERTIES CXX_SCAN_FOR_MODULES OFF)
if(EMSCRIPTEN)
set(EMDAWNWEBGPU_DIR "" CACHE PATH "Path to emdawnwebgpu_pkg")
+26 -17
View File
@@ -374,20 +374,28 @@ static wgpu::Buffer ggml_webgpu_tensor_buf(const ggml_tensor * tensor) {
return ctx->buffer;
}
// Binding offset for a tensor: the largest aligned offset at or before the tensor whose
// distance to the tensor is a whole number of type blocks, so shaders can index the
// misalignment in elements even for block quantized types.
static size_t ggml_webgpu_tensor_align_offset(const ggml_tensor * t, size_t alignment) {
const size_t offset = ggml_webgpu_tensor_offset(t);
const size_t type_size = ggml_type_size(t->type);
size_t aligned = offset & ~(alignment - 1);
while ((offset - aligned) % type_size != 0) {
GGML_ASSERT(aligned >= alignment);
aligned -= alignment;
}
return aligned;
}
static size_t ggml_webgpu_tensor_misalignment(const ggml_tensor * t, size_t alignment) {
size_t offset = ggml_webgpu_tensor_offset(t);
return offset & (alignment - 1);
return ggml_webgpu_tensor_offset(t) - ggml_webgpu_tensor_align_offset(t, alignment);
}
static size_t ggml_webgpu_tensor_misalignment(webgpu_context & ctx, const ggml_tensor * t) {
return ggml_webgpu_tensor_misalignment(t, ctx->global_ctx->capabilities.limits.minStorageBufferOffsetAlignment);
}
static size_t ggml_webgpu_tensor_align_offset(const ggml_tensor * t, size_t alignment) {
size_t offset = ggml_webgpu_tensor_offset(t);
return offset & ~(alignment - 1);
}
static size_t ggml_webgpu_tensor_align_offset(webgpu_context & ctx, const ggml_tensor * t) {
return ggml_webgpu_tensor_align_offset(t, ctx->global_ctx->capabilities.limits.minStorageBufferOffsetAlignment);
}
@@ -4006,16 +4014,17 @@ static void ggml_backend_webgpu_request_adapter(wgpu::Instance & instance, wgpu:
options.nextInChain = &adapterTogglesDesc;
#endif
instance.WaitAny(instance.RequestAdapter(
&options, wgpu::CallbackMode::AllowSpontaneous,
[&adapter](wgpu::RequestAdapterStatus status, wgpu::Adapter _adapter, const char * message) {
if (status != wgpu::RequestAdapterStatus::Success) {
GGML_LOG_ERROR("ggml_webgpu: Failed to get an adapter: %s\n", message);
return;
}
adapter = std::move(_adapter);
}),
UINT64_MAX);
instance.WaitAny(
instance.RequestAdapter(
&options, wgpu::CallbackMode::AllowSpontaneous,
[&adapter](wgpu::RequestAdapterStatus status, wgpu::Adapter _adapter, wgpu::StringView message) {
if (status != wgpu::RequestAdapterStatus::Success) {
GGML_LOG_ERROR("ggml_webgpu: Failed to get an adapter: %s\n", std::string(message).c_str());
return;
}
adapter = std::move(_adapter);
}),
UINT64_MAX);
}
static void create_webgpu_device(ggml_backend_webgpu_reg_context * ctx) {
+136
View File
@@ -0,0 +1,136 @@
# Compile-time profiling using clang -ftime-trace + ClangBuildAnalyzer.
#
# Usage:
# .\scripts\build-profile.ps1 [-Full] [-Jobs N]
#
# -Full : include Server, Tools, and Tests (default: minimal build)
# -Jobs : number of parallel jobs (default: all cores)
#
# Requires ClangBuildAnalyzer:
# https://github.com/aras-p/ClangBuildAnalyzer
param(
[switch]$Full,
[int]$Jobs = [Environment]::ProcessorCount
)
$ErrorActionPreference = "Stop"
$ScriptDir = Split-Path -Parent $MyInvocation.MyCommand.Path
$RootDir = Split-Path -Parent $ScriptDir
if ($Full) {
$BuildDir = Join-Path $RootDir "build-profile-full"
$Report = Join-Path $BuildDir "profile-report-full.txt"
} else {
$BuildDir = Join-Path $RootDir "build-profile-baseline"
$Report = Join-Path $BuildDir "profile-report.txt"
}
$OutputBin = Join-Path $BuildDir "clang_analysis.bin"
if (-not (Get-Command clang++ -ErrorAction SilentlyContinue)) {
Write-Error "clang++ not found"
exit 1
}
if (-not (Get-Command ninja -ErrorAction SilentlyContinue)) {
Write-Error "ninja not found (required so cmake does not fall back to the Visual Studio/MSVC generator)"
exit 1
}
if (-not (Get-Command ClangBuildAnalyzer -ErrorAction SilentlyContinue)) {
Write-Error "ClangBuildAnalyzer not found`n https://github.com/aras-p/ClangBuildAnalyzer/releases"
exit 1
}
$ClangVer = (clang++ --version | Select-Object -First 1)
Write-Host "compiler : $ClangVer"
Write-Host "build dir: $BuildDir"
Write-Host "output : $OutputBin"
Write-Host "jobs : $Jobs"
Write-Host ""
if (Get-Command ccache -ErrorAction SilentlyContinue) {
Write-Host "clearing ccache..."
ccache -C -z
}
$env:CCACHE_DISABLE = "1"
$TestsFlag = if ($Full) { "ON" } else { "OFF" }
$ToolsFlag = if ($Full) { "ON" } else { "OFF" }
$ServerFlag = if ($Full) { "ON" } else { "OFF" }
cmake --fresh `
-S $RootDir `
-B $BuildDir `
-G "Ninja" `
-DCMAKE_BUILD_TYPE=Release `
-DCMAKE_C_COMPILER=clang `
-DCMAKE_CXX_COMPILER=clang++ `
-DCMAKE_C_FLAGS="-ftime-trace" `
-DCMAKE_CXX_FLAGS="-ftime-trace" `
-DGGML_CCACHE=OFF `
-DGGML_OPENMP=ON `
-DGGML_NATIVE=OFF `
"-DLLAMA_BUILD_TESTS=$TestsFlag" `
-DLLAMA_BUILD_EXAMPLES=OFF `
"-DLLAMA_BUILD_TOOLS=$ToolsFlag" `
"-DLLAMA_BUILD_SERVER=$ServerFlag" `
-DLLAMA_BUILD_APP=OFF
if ($LASTEXITCODE -ne 0) { exit $LASTEXITCODE }
$StrayTrace = Join-Path $RootDir "-.json"
if (Test-Path $StrayTrace) {
Remove-Item $StrayTrace -Force
}
Write-Host ""
Write-Host "Initializing ClangBuildAnalyzer..."
ClangBuildAnalyzer --start $BuildDir
Write-Host ""
Write-Host "building..."
Write-Host ""
$StartTime = Get-Date
cmake --build $BuildDir --clean-first -j $Jobs
if ($LASTEXITCODE -ne 0) { exit $LASTEXITCODE }
$Elapsed = (Get-Date) - $StartTime
Write-Host ""
Write-Host ("build time: {0}s ({1}m {2}s)" -f [int]$Elapsed.TotalSeconds, [int]$Elapsed.TotalMinutes, $Elapsed.Seconds)
Write-Host ""
Write-Host "Aggregating profile metrics..."
ClangBuildAnalyzer --stop $BuildDir $OutputBin | Out-Null
Write-Host ""
Write-Host ("=" * 80)
$TUs = "?"
if (Test-Path $Report) {
$Match = Select-String -Path $Report -Pattern "Compilation \((\d+)" | Select-Object -First 1
if ($Match) { $TUs = $Match.Matches[0].Groups[1].Value }
}
ClangBuildAnalyzer --analyze $OutputBin | Tee-Object -FilePath $Report
Write-Host ""
Write-Host "translation units: $TUs"
Write-Host ""
Write-Host "largest trace files (top 20 by size):"
Get-ChildItem -Path $BuildDir -Recurse -Filter "*.json" |
Where-Object { $_.Name -ne "compile_commands.json" } |
Sort-Object Length -Descending |
Select-Object -First 20 |
ForEach-Object { "{0,8:F1} KB {1}" -f ($_.Length / 1024), $_.FullName }
Write-Host ""
Write-Host "ClangBuildAnalyzer report was generated: $Report"
+122
View File
@@ -0,0 +1,122 @@
#!/usr/bin/env bash
# Compile-time profiling using clang -ftime-trace + ClangBuildAnalyzer.
#
# Usage:
# ./scripts/build-profile.sh [--full] [-jN]
#
# --full: include Server, Tools, and Tests (default: minimal build)
# -jN : number of parallel jobs (default: all cores)
#
# Requires ClangBuildAnalyzer:
# macOS: brew install clang-build-analyzer
# Linux: https://github.com/aras-p/ClangBuildAnalyzer.git
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
ROOT_DIR="$(cd "${SCRIPT_DIR}/.." && pwd)"
FULL=0
JOBS="-j$(nproc 2>/dev/null || sysctl -n hw.ncpu)"
for arg in "$@"; do
case "${arg}" in
--full) FULL=1 ;;
-j*) JOBS="${arg}" ;;
*) echo "error: unknown argument: ${arg}" >&2; exit 1 ;;
esac
done
if [ "${FULL}" -eq 1 ]; then
BUILD_DIR="${ROOT_DIR}/build-profile-full"
REPORT="${BUILD_DIR}/profile-report-full.txt"
else
BUILD_DIR="${ROOT_DIR}/build-profile-baseline"
REPORT="${BUILD_DIR}/profile-report.txt"
fi
OUTPUT_BIN="${BUILD_DIR}/clang_analysis.bin"
if ! command -v clang++ &>/dev/null; then
echo "error: clang++ not found" >&2
exit 1
fi
if ! command -v ClangBuildAnalyzer &>/dev/null; then
echo "error: ClangBuildAnalyzer not found" >&2
echo " brew install clangbuildanalyzer (macOS)" >&2
echo " or: https://github.com/aras-p/ClangBuildAnalyzer/releases" >&2
exit 1
fi
CLANG_VER=$(clang++ --version | head -1)
echo "compiler : ${CLANG_VER}"
echo "build dir: ${BUILD_DIR}"
echo "output : ${OUTPUT_BIN}"
echo "jobs : ${JOBS}"
echo
if command -v ccache &>/dev/null; then
echo "clearing ccache..."
ccache -C -z
fi
export CCACHE_DISABLE=1
cmake --fresh \
-S "${ROOT_DIR}" \
-B "${BUILD_DIR}" \
-DCMAKE_BUILD_TYPE=Release \
-DCMAKE_C_COMPILER=clang \
-DCMAKE_CXX_COMPILER=clang++ \
-DCMAKE_C_FLAGS="-ftime-trace" \
-DCMAKE_CXX_FLAGS="-ftime-trace" \
-DGGML_CCACHE=OFF \
-DGGML_OPENMP=ON \
-DGGML_NATIVE=OFF \
-DLLAMA_BUILD_TESTS=$([ "${FULL}" -eq 1 ] && echo ON || echo OFF) \
-DLLAMA_BUILD_EXAMPLES=OFF \
-DLLAMA_BUILD_TOOLS=$([ "${FULL}" -eq 1 ] && echo ON || echo OFF) \
-DLLAMA_BUILD_SERVER=$([ "${FULL}" -eq 1 ] && echo ON || echo OFF) \
-DLLAMA_BUILD_APP=OFF
echo
echo "Initializing ClangBuildAnalyzer..."
ClangBuildAnalyzer --start "${BUILD_DIR}"
echo
echo "building..."
echo
START=$(date +%s)
cmake --build "${BUILD_DIR}" --clean-first "${JOBS}"
END=$(date +%s)
ELAPSED=$((END - START))
echo
printf "build time: %ds (%dm %ds)\n" "${ELAPSED}" "$((ELAPSED / 60))" "$((ELAPSED % 60))"
echo
echo "Aggregating profile metrics..."
ClangBuildAnalyzer --stop "${BUILD_DIR}" "${OUTPUT_BIN}" > /dev/null
echo
echo "================================================================================"
TUS=$(grep -oP "Compilation \(\K[0-9]+" "${REPORT}" 2>/dev/null || echo "?")
ClangBuildAnalyzer --analyze "${OUTPUT_BIN}" | tee "${REPORT}"
echo
echo "translation units: ${TUS}"
echo
echo "largest trace files (top 20 by size):"
find "${BUILD_DIR}" -name "*.json" ! -name "compile_commands.json" \
| xargs ls -l 2>/dev/null \
| awk 'NF>5 {print $5, $NF}' \
| sort -rn \
| awk 'NR<=20 {printf "%8.1f KB %s\n", $1/1024, $2}'
echo
echo "ClangBuildAnalyzer report was generated: ${REPORT}"
+296
View File
@@ -0,0 +1,296 @@
#!/usr/bin/env python3
"""
align-macros.py - Inspect and align trailing backslashes in multiline C/C++ macros.
Usage:
align-macros.py [paths...] # Check and report misaligned macros
align-macros.py --diff [paths...] # Show unified diff of fixes
align-macros.py --fix [paths...] # Fix misaligned macros in-place
align-macros.py --fix --mode majority ... # Align to the dominant column
align-macros.py --fix --pad 2 ... # Align to (max_content_len + pad)
Safety rules:
- Macros that are ALREADY aligned are NEVER touched (unless --all is given).
- Whitespace after trailing backslashes is flagged and cleaned.
"""
import argparse
import difflib
import logging
import os
import re
import sys
from collections import Counter
from typing import List, Optional, Tuple, NamedTuple
logger = logging.getLogger("ggml-hexagon-align-macros")
class MacroLine(NamedTuple):
line_num: int # 1-indexed
raw: str # Original line including newline
content: str # Line content before trailing backslash (stripped of trailing whitespace)
bs_col: Optional[int] # 1-indexed column of backslash, or None if last line has no backslash
trailing_ws: bool # True if whitespace existed after the backslash
class MacroDef(NamedTuple):
name: str
filepath: str
start_line: int
end_line: int
lines: List[MacroLine]
def parse_macros(filepath: str) -> List[MacroDef]:
"""Extract all multiline macros from a C/C++ source file."""
try:
with open(filepath, "r", encoding="utf-8", errors="replace") as f:
lines = f.readlines()
except Exception as e:
logger.error(f"Error reading {filepath}: {e}")
return []
macros: List[MacroDef] = []
i = 0
n = len(lines)
while i < n:
line = lines[i]
m = re.match(r"^\s*#\s*define\s+([A-Za-z_][A-Za-z0-9_]*)", line)
if m:
macro_name = m.group(1)
macro_start = i + 1
macro_lines: List[MacroLine] = []
cur = i
while cur < n:
l_raw = lines[cur]
l_rstrip = l_raw.rstrip("\r\n")
# Check if line has a trailing backslash
# Note: handle possible accidental spaces after backslash
match_bs = re.search(r"\\([ \t]*)$", l_rstrip)
if match_bs:
has_trailing_ws = len(match_bs.group(1)) > 0
bs_index = match_bs.start()
content = l_rstrip[:bs_index].rstrip()
# 1-indexed column of the backslash
bs_col = bs_index + 1
macro_lines.append(MacroLine(
line_num=cur + 1,
raw=l_raw,
content=content,
bs_col=bs_col,
trailing_ws=has_trailing_ws
))
cur += 1
else:
# Line does not end with backslash
if cur == i:
# Single-line macro, not multiline
break
else:
# Final line of a multiline macro
macro_lines.append(MacroLine(
line_num=cur + 1,
raw=l_raw,
content=l_rstrip.rstrip(),
bs_col=None,
trailing_ws=False
))
break
# Only record if it is a multiline macro (has at least one continuation line)
continuation_lines = [ml for ml in macro_lines if ml.bs_col is not None]
if continuation_lines:
macro_end = macro_lines[-1].line_num
macros.append(MacroDef(
name=macro_name,
filepath=filepath,
start_line=macro_start,
end_line=macro_end,
lines=macro_lines
))
i = cur
i += 1
return macros
def is_macro_aligned(macro: MacroDef) -> bool:
"""A macro is aligned if all continuation lines have backslashes at the same column."""
bs_cols = [ml.bs_col for ml in macro.lines if ml.bs_col is not None]
if not bs_cols:
return True
has_trailing_ws = any(ml.trailing_ws for ml in macro.lines)
return len(set(bs_cols)) == 1 and not has_trailing_ws
def compute_target_column(macro: MacroDef, mode: str, pad: int, target_col: Optional[int]) -> int:
"""Determine the column where backslashes should be aligned."""
max_content_len = max(len(ml.content) for ml in macro.lines)
min_needed = max_content_len + pad
if target_col is not None:
return max(target_col, min_needed)
bs_cols = [ml.bs_col for ml in macro.lines if ml.bs_col is not None]
if not bs_cols:
return min_needed
if mode == "min":
return min_needed
elif mode == "max":
return max(max(bs_cols), min_needed)
elif mode == "majority":
counts = Counter(bs_cols)
# Sort by frequency descending, then by column descending
majority_col = sorted(counts.items(), key=lambda x: (-x[1], -x[0]))[0][0]
return max(majority_col, min_needed)
else:
return min_needed
def realign_macro_lines(macro: MacroDef, target_col: int) -> List[str]:
"""Format macro lines with backslashes aligned at target_col."""
new_lines: List[str] = []
for ml in macro.lines:
nl = "\r\n" if ml.raw.endswith("\r\n") else "\n"
if ml.bs_col is None:
# Last line without backslash
new_lines.append(ml.raw)
else:
if not ml.content:
spaces = " " * (target_col - 1)
new_lines.append(f"{spaces}\\{nl}")
else:
spaces_needed = max(1, target_col - len(ml.content) - 1)
new_lines.append(f"{ml.content}{' ' * spaces_needed}\\{nl}")
return new_lines
def process_file(filepath: str, args: argparse.Namespace) -> Tuple[int, int, Optional[str]]:
macros = parse_macros(filepath)
if not macros:
return 0, 0, None
with open(filepath, "r", encoding="utf-8", errors="replace") as f:
file_lines = f.readlines()
misaligned_count = 0
modified = False
new_file_lines = list(file_lines)
for macro in macros:
aligned = is_macro_aligned(macro)
if not aligned or args.all:
if not aligned:
misaligned_count += 1
bs_cols = [ml.bs_col for ml in macro.lines if ml.bs_col is not None]
max_content = max(len(ml.content) for ml in macro.lines)
col_counts = Counter(bs_cols)
if not args.quiet:
logger.info(f"{filepath}:{macro.start_line}-{macro.end_line} [{macro.name}]")
logger.info(f" Max content width: {max_content}, Min needed column (+{args.pad}): {max_content + args.pad}")
logger.info(f" Current backslash columns: {dict(sorted(col_counts.items()))}")
trailing_ws_lines = [ml.line_num for ml in macro.lines if ml.trailing_ws]
if trailing_ws_lines:
logger.warning(f" Warning: Trailing whitespace after backslash on line(s): {trailing_ws_lines}")
target_col = compute_target_column(macro, args.mode, args.pad, args.target_col)
if not args.quiet:
logger.info(f" -> Target alignment column: {target_col}")
realigned = realign_macro_lines(macro, target_col)
start_idx = macro.start_line - 1
end_idx = start_idx + len(macro.lines)
if new_file_lines[start_idx:end_idx] != realigned:
new_file_lines[start_idx:end_idx] = realigned
modified = True
diff_text = None
if modified:
diff = difflib.unified_diff(
file_lines,
new_file_lines,
fromfile=f"a/{filepath}",
tofile=f"b/{filepath}",
lineterm=""
)
diff_text = "\n".join(diff)
if args.fix:
with open(filepath, "w", encoding="utf-8") as f:
f.writelines(new_file_lines)
if not args.quiet:
logger.info(f" [FIXED] Updated {filepath}")
return len(macros), misaligned_count, diff_text
def find_source_files(paths: List[str]) -> List[str]:
extensions = {".c", ".cpp", ".cc", ".cxx", ".h", ".hpp", ".inl"}
result: List[str] = []
for p in paths:
if os.path.isfile(p):
result.append(p)
elif os.path.isdir(p):
for root, _, files in os.walk(p):
for file in sorted(files):
_, ext = os.path.splitext(file)
if ext.lower() in extensions:
result.append(os.path.join(root, file))
return sorted(result)
def main():
logging.basicConfig(level=logging.INFO, format="%(message)s")
parser = argparse.ArgumentParser(
description="Inspect and align backslashes in multiline C/C++ macros."
)
parser.add_argument("paths", nargs="*", default=["."], help="Files or directories to scan (default: current dir)")
parser.add_argument("--fix", action="store_true", help="Fix misaligned macros in-place")
parser.add_argument("--diff", action="store_true", help="Display unified diff of suggested fixes")
parser.add_argument("--check", action="store_true", help="Exit with code 1 if misaligned macros exist")
parser.add_argument("--mode", choices=["min", "max", "majority"], default="min",
help="Alignment mode: 'min' (max_len + pad), 'max' (max existing col), 'majority' (dominant col)")
parser.add_argument("--pad", type=int, default=2, help="Spaces between longest line and backslash (default: 2)")
parser.add_argument("--target-col", type=int, default=None, help="Force alignment to an exact column")
parser.add_argument("--all", action="store_true", help="Realign all macros even if already aligned (default: only misaligned)")
parser.add_argument("-q", "--quiet", action="store_true", help="Only output errors and diffs/summary")
args = parser.parse_args()
files = find_source_files(args.paths)
if not files:
logger.error("No C/C++ source files found.")
sys.exit(0)
total_macros = 0
total_misaligned = 0
diffs: List[str] = []
for filepath in files:
num_macros, num_misaligned, diff_text = process_file(filepath, args)
total_macros += num_macros
total_misaligned += num_misaligned
if diff_text:
diffs.append(diff_text)
if args.diff and diffs:
logger.info("\n--- Proposed Changes ---\n")
for d in diffs:
logger.info(d)
logger.info(f"\nSummary: scanned {len(files)} files, {total_macros} multiline macros, {total_misaligned} misaligned.")
if args.check and total_misaligned > 0:
sys.exit(1)
if __name__ == "__main__":
main()
+95 -33
View File
@@ -14,6 +14,42 @@ import logging
logger = logging.getLogger("run")
MANAGED_ENV_NAMES = (
"GGML_HEXAGON_DEVICES",
"GGML_HEXAGON_VERBOSE",
"GGML_HEXAGON_PROFILE",
"GGML_HEXAGON_NHVX",
"GGML_HEXAGON_NHMX",
"GGML_HEXAGON_HOSTBUF",
"GGML_HEXAGON_OPBATCH",
"GGML_HEXAGON_OPQUEUE",
"GGML_HEXAGON_OPPOLL",
"GGML_HEXAGON_OPFILTER",
"GGML_HEXAGON_OPFUSION",
"GGML_HEXAGON_VMEM",
"GGML_HEXAGON_MBUF",
"GGML_HEXAGON_MM_SELECT",
"GGML_HEXAGON_FA_SELECT",
"GGML_HEXAGON_AR_SELECT",
"GGML_HEXAGON_ETM",
"GGML_HEXAGON_ARCH",
"GGML_HEXAGON_OPTRACE",
"GGML_OPENCL_PLATFORM",
"GGML_OPENCL_DEVICE",
"GGML_OPENCL_OPFILTER",
"GGML_OPENCL_KERNEL_CACHE_DIR",
"GGML_OPENCL_KERNEL_CACHE_DEBUG",
"GGML_OPENCL_FA_TUNE",
"GGML_OPENCL_DISABLE_FUSION",
"GGML_OPENCL_ADRENO_XMEM_GEMM",
"GGML_OPENCL_ADRENO_USE_LARGE_BUFFER",
"GGML_SCHED_DEBUG",
"MTMD_BACKEND_DEVICE",
"D",
"DEVICE",
)
def parse_target(target_str):
if not target_str:
return None, None
@@ -38,6 +74,57 @@ def shlex_join(args_list):
return " ".join(pipes.quote(x) for x in args_list)
def split_device_list(devices):
parts = []
curr = []
bracket_depth = 0
for ch in devices:
if ch == '[':
bracket_depth += 1
curr.append(ch)
elif ch == ']':
if bracket_depth > 0:
bracket_depth -= 1
curr.append(ch)
elif ch == ',' and bracket_depth == 0:
part = "".join(curr).strip()
if part:
parts.append(part)
curr = []
else:
curr.append(ch)
part = "".join(curr).strip()
if part:
parts.append(part)
return parts
def device_arg_from_devices(devices):
if devices.isdigit():
n = int(devices)
return ",".join(f"HTP{i}" for i in range(n))
names = []
for part in split_device_list(devices):
if "[" in part:
part = part.split("[", 1)[0].strip()
if part:
names.append(part)
return ",".join(names)
def normalize_cmd_device_args(cmd_args):
for i, arg in enumerate(cmd_args):
if arg == "--device" and i + 1 < len(cmd_args):
cmd_args[i + 1] = device_arg_from_devices(cmd_args[i + 1])
elif arg.startswith("--device="):
cmd_args[i] = "--device=" + device_arg_from_devices(arg.split("=", 1)[1])
def main():
logging.basicConfig(level=logging.INFO, format='%(message)s')
# Split arguments at '--'
@@ -142,8 +229,6 @@ def main():
def set_env(env_name, opt_val):
if opt_val is not None:
env_vars[env_name] = str(opt_val)
elif env_name in os.environ:
env_vars[env_name] = os.environ[env_name]
# Resolve and filter devices (HTP vs OpenCL)
device_in_cmd = None
@@ -166,7 +251,7 @@ def main():
hex_devices = devices_val
cl_device = ""
else:
parts = [p.strip() for p in devices_val.split(",")]
parts = split_device_list(devices_val)
# Any device containing "htp" is Hexagon, rest is OpenCL
hex_parts = [p for p in parts if "htp" in p.lower()]
cl_parts = [
@@ -181,15 +266,13 @@ def main():
# Set Hexagon devices
if hex_devices:
env_vars["GGML_HEXAGON_DEVICES"] = hex_devices
elif "GGML_HEXAGON_DEVICES" in os.environ:
env_vars["GGML_HEXAGON_DEVICES"] = os.environ["GGML_HEXAGON_DEVICES"]
normalize_cmd_device_args(cmd_args)
# Set OpenCL device (unless overridden by --cl-device)
final_cl_device = args.cl_device if args.cl_device is not None else cl_device
if final_cl_device:
env_vars["GGML_OPENCL_DEVICE"] = final_cl_device
elif "GGML_OPENCL_DEVICE" in os.environ:
env_vars["GGML_OPENCL_DEVICE"] = os.environ["GGML_OPENCL_DEVICE"]
# Map shared & backend-specific parameters with correct overrides
@@ -206,8 +289,6 @@ def main():
if args.cl_fa_tune or args.profile is not None:
env_vars["GGML_OPENCL_FA_TUNE"] = "1"
elif "GGML_OPENCL_FA_TUNE" in os.environ:
env_vars["GGML_OPENCL_FA_TUNE"] = os.environ["GGML_OPENCL_FA_TUNE"]
# Other Hexagon environment variables
set_env("GGML_HEXAGON_NHVX", args.hex_nhvx)
@@ -235,18 +316,12 @@ def main():
if args.cl_disable_fusion:
env_vars["GGML_OPENCL_DISABLE_FUSION"] = "1"
elif "GGML_OPENCL_DISABLE_FUSION" in os.environ:
env_vars["GGML_OPENCL_DISABLE_FUSION"] = os.environ["GGML_OPENCL_DISABLE_FUSION"]
if args.cl_adreno_xmem:
env_vars["GGML_OPENCL_ADRENO_XMEM_GEMM"] = "1"
elif "GGML_OPENCL_ADRENO_XMEM_GEMM" in os.environ:
env_vars["GGML_OPENCL_ADRENO_XMEM_GEMM"] = os.environ["GGML_OPENCL_ADRENO_XMEM_GEMM"]
if args.cl_adreno_large_buffer:
env_vars["GGML_OPENCL_ADRENO_USE_LARGE_BUFFER"] = "1"
elif "GGML_OPENCL_ADRENO_USE_LARGE_BUFFER" in os.environ:
env_vars["GGML_OPENCL_ADRENO_USE_LARGE_BUFFER"] = os.environ["GGML_OPENCL_ADRENO_USE_LARGE_BUFFER"]
if args.sched_debug:
env_vars["GGML_SCHED_DEBUG"] = "2"
@@ -288,15 +363,7 @@ def main():
has_b = any(arg == "-b" for arg in cmd_args)
if not has_b:
if args.devices:
if args.devices.isdigit():
n = int(args.devices)
device_val = ",".join(f"HTP{i}" for i in range(n))
else:
device_val = args.devices
elif "D" in os.environ:
device_val = os.environ["D"]
elif "DEVICE" in os.environ:
device_val = os.environ["DEVICE"]
device_val = device_arg_from_devices(args.devices)
else:
device_val = "HTP0"
if device_val:
@@ -305,17 +372,10 @@ def main():
has_device = any(arg.startswith("--device") for arg in cmd_args)
if not has_device:
if args.devices:
if args.devices.isdigit():
n = int(args.devices)
device_val = ",".join(f"HTP{i}" for i in range(n))
else:
device_val = args.devices
elif "D" in os.environ:
device_val = os.environ["D"]
elif "DEVICE" in os.environ:
device_val = os.environ["DEVICE"]
device_val = device_arg_from_devices(args.devices)
else:
device_val = "HTP0"
if device_val:
cmd_args += ["--device", device_val]
@@ -415,6 +475,8 @@ def main():
else:
local_env["LD_LIBRARY_PATH"] = lib_dir + os.path.pathsep + local_env.get("LD_LIBRARY_PATH", "")
for k in MANAGED_ENV_NAMES:
local_env.pop(k, None)
for k, v in env_vars.items():
local_env[k] = v
+1 -1
View File
@@ -5,7 +5,7 @@ import os
import sys
import subprocess
HTTPLIB_VERSION = "refs/tags/v0.54.1"
HTTPLIB_VERSION = "refs/tags/v0.56.0"
# used by examples/gguf-hash, these repos have no release tag, so we pin a commit
XXHASH_COMMIT = "9f465f1ea932d6ad9a26cd77496311ffa544cd68"
+104 -53
View File
@@ -21,6 +21,7 @@ set(DIST_DIR "${UI_BINARY_DIR}/dist")
set(SRC_DIST_DIR "${UI_SOURCE_DIR}/dist")
set(WORK_DIR "${UI_BINARY_DIR}/ui-src")
set(STAMP_FILE "${UI_BINARY_DIR}/.ui-stamp")
set(EMBED_STAMP "${UI_BINARY_DIR}/.ui-embed.sha256")
set(UI_CPP "${UI_BINARY_DIR}/ui.cpp")
set(UI_H "${UI_BINARY_DIR}/ui.h")
@@ -141,9 +142,59 @@ function(ui_validate_assets files in_dir)
endfunction()
# Generate ui.cpp/ui.h embedding every file of ${dist_dir} (empty table when
# it has no index.html). When LLAMA_UI_GZIP is enabled, assets are compressed
# first and served pre-gzipped (llama_ui_use_gzip()).
# it has no index.html), gzip-compressed when LLAMA_UI_GZIP is enabled.
function(emit_files dist_dir)
set(UI_TEMPLATE_DIR "${LLAMA_SOURCE_DIR}/tools/ui")
# Collect the asset list once and reuse it for the fingerprint,
# validation, compression and embedding.
set(assets "")
if(EXISTS "${dist_dir}/index.html")
file(GLOB_RECURSE assets
LIST_DIRECTORIES false
RELATIVE "${dist_dir}"
"${dist_dir}/*")
list(FILTER assets EXCLUDE REGEX "^_gzip/")
list(SORT assets)
endif()
if(LLAMA_UI_GZIP AND NOT DEFINED ENV{SOURCE_DATE_EPOCH})
# Zero the gzip header timestamp so identical inputs give identical
# bytes (and therefore stable ETags) on every machine.
set(ENV{SOURCE_DATE_EPOCH} 0)
endif()
# Fingerprint of every input that determines ui.cpp/ui.h: compression
# settings, the asset tree (names + SHA-256) and this script + templates.
set(fp "${LLAMA_UI_GZIP}|$ENV{SOURCE_DATE_EPOCH}|${CMAKE_VERSION}\n")
foreach(f ${assets})
file(SHA256 "${dist_dir}/${f}" h)
string(APPEND fp "${f} ${h}\n")
endforeach()
foreach(g
"${CMAKE_CURRENT_FUNCTION_LIST_FILE}"
"${UI_TEMPLATE_DIR}/ui.h.in"
"${UI_TEMPLATE_DIR}/ui.cpp.in")
file(SHA256 "${g}" h)
string(APPEND fp "gen ${h}\n")
endforeach()
string(SHA256 fingerprint "${fp}")
if(EXISTS "${EMBED_STAMP}" AND EXISTS "${UI_CPP}" AND EXISTS "${UI_H}")
file(READ "${EMBED_STAMP}" fp_saved)
string(STRIP "${fp_saved}" fp_saved)
if(fp_saved STREQUAL "${fingerprint}")
message(STATUS "UI: assets unchanged, skipping embedding")
return()
endif()
endif()
# Drop the old stamp up front so a crash mid-generation cannot leave
# outputs and stamp out of sync.
file(REMOVE "${EMBED_STAMP}")
ui_validate_assets("${assets}" "${dist_dir}")
set(embed_dir "${dist_dir}")
set(use_gzip FALSE)
@@ -156,21 +207,11 @@ function(emit_files dist_dir)
endif()
if(LLAMA_UI_GZIP)
# Compress every asset into a parallel _gzip/ tree under the build
# directory (never write into the source or dist tree); the
# structure stays the same: /abc/def --> /_gzip/abc/def.
# FORMAT raw produces a bare gzip stream (no archive container)
# that can be served with Content-Encoding: gzip. SOURCE_DATE_EPOCH
# zeroes the header timestamp so identical inputs give identical
# bytes (and therefore stable ETags) on every machine.
if(NOT DEFINED ENV{SOURCE_DATE_EPOCH})
set(ENV{SOURCE_DATE_EPOCH} 0)
endif()
# directory, served with Content-Encoding: gzip.
set(gzip_root "${UI_BINARY_DIR}/ui-gzip")
set(gzip_dir "${gzip_root}/_gzip")
file(REMOVE_RECURSE "${gzip_root}")
file(GLOB_RECURSE all_files RELATIVE "${dist_dir}" "${dist_dir}/*")
list(FILTER all_files EXCLUDE REGEX "^_gzip/")
foreach(f ${all_files})
foreach(f IN LISTS assets)
get_filename_component(asset_path "${dist_dir}/${f}" REALPATH)
get_filename_component(dst_dir "${gzip_dir}/${f}" DIRECTORY)
file(MAKE_DIRECTORY "${dst_dir}")
@@ -187,21 +228,10 @@ function(emit_files dist_dir)
endif()
endif()
set(assets "")
if(EXISTS "${embed_dir}/index.html")
file(GLOB_RECURSE assets RELATIVE "${embed_dir}" "${embed_dir}/*")
list(FILTER assets EXCLUDE REGEX "^_gzip/")
list(SORT assets)
ui_validate_assets("${assets}" "${embed_dir}")
endif()
list(LENGTH assets n_assets)
# Only the per-asset data arrays and table rows are built here; all
# static C++ lives in the ui.h.in / ui.cpp.in templates. configure_file
# rewrites an output only when its contents change, so the library is
# not recompiled needlessly. @ONLY keeps ${...} in the content literal;
# mime types come from a fixed list.
# Per-asset arrays and table rows go into the ui.h.in / ui.cpp.in templates;
# configure_file only rewrites on content change, avoiding needless recompiles.
set(ASSET_ARRAYS "")
set(ASSET_TABLE "")
set(idx 0)
@@ -235,9 +265,11 @@ function(emit_files dist_dir)
set(USE_GZIP true)
endif()
set(UI_TEMPLATE_DIR "${LLAMA_SOURCE_DIR}/tools/ui")
configure_file("${UI_TEMPLATE_DIR}/ui.h.in" "${UI_H}" @ONLY)
configure_file("${UI_TEMPLATE_DIR}/ui.cpp.in" "${UI_CPP}" @ONLY)
# Write the embed stamp last, after both generated files succeeded.
file(WRITE "${EMBED_STAMP}" "${fingerprint}")
message(STATUS "UI: embedded ${n_assets} assets")
endfunction()
@@ -419,16 +451,8 @@ function(hf_download version out_var out_resolved)
message(STATUS "UI: downloading from ${resolved}: ${base}/dist.tar.gz")
file(DOWNLOAD "${base}/dist.tar.gz?download=true" "${archive}"
STATUS status TIMEOUT 300 ${auth_headers}
)
list(GET status 0 rc)
if(NOT rc EQUAL 0)
list(GET status 1 errmsg)
message(STATUS "UI: download dist.tar.gz from ${resolved} failed: ${errmsg}")
continue()
endif()
# Fetch the checksum first: when the archive we already have matches
# it, the expensive download is skipped and only extraction repeats.
file(DOWNLOAD "${base}/dist.tar.gz.sha256?download=true" "${archive}.sha256"
STATUS status TIMEOUT 30 ${auth_headers}
)
@@ -439,17 +463,44 @@ function(hf_download version out_var out_resolved)
continue()
endif()
# Validate sha256 checkums
# Validate the sha256 checksum: reject anything that is not a full
# 64-hex-digit digest before touching the archive.
file(READ "${archive}.sha256" expected)
string(REGEX MATCH "^[0-9a-fA-F]+" expected "${expected}")
string(TOLOWER "${expected}" expected)
file(SHA256 "${archive}" actual)
if("${expected}" STREQUAL "" OR NOT "${actual}" STREQUAL "${expected}")
message(STATUS "UI: checksum mismatch for dist.tar.gz from ${resolved}")
string(LENGTH "${expected}" expected_len)
if(NOT expected_len EQUAL 64)
message(STATUS "UI: invalid checksum from ${resolved}")
continue()
endif()
# Clear DIST_DIR to remove stale files first
set(actual "")
if(EXISTS "${archive}")
file(SHA256 "${archive}" actual)
endif()
if("${actual}" STREQUAL "${expected}")
message(STATUS "UI: local dist.tar.gz matches checksum from ${resolved}, skipping download")
else()
file(DOWNLOAD "${base}/dist.tar.gz?download=true" "${archive}"
STATUS status TIMEOUT 300 ${auth_headers}
)
list(GET status 0 rc)
if(NOT rc EQUAL 0)
list(GET status 1 errmsg)
message(STATUS "UI: download dist.tar.gz from ${resolved} failed: ${errmsg}")
continue()
endif()
file(SHA256 "${archive}" actual)
if(NOT "${actual}" STREQUAL "${expected}")
message(STATUS "UI: checksum mismatch for dist.tar.gz from ${resolved}")
continue()
endif()
endif()
# Remove the stamp with the dist tree it describes, together.
file(REMOVE "${STAMP_FILE}")
file(REMOVE_RECURSE "${DIST_DIR}")
file(ARCHIVE_EXTRACT INPUT "${archive}" DESTINATION "${DIST_DIR}")
@@ -495,27 +546,27 @@ endif()
if(NOT provisioned AND HF_ENABLED)
resolve_version(VERSION)
# Stamp a successful HF download: records bucket + requested version and
# lets later steps distinguish downloaded assets from locally built ones.
set(stamp_key "${HF_BUCKET}|${VERSION}")
set(stamp_ok FALSE)
if(EXISTS "${STAMP_FILE}" AND NOT "${VERSION}" STREQUAL "")
if(EXISTS "${STAMP_FILE}" AND EXISTS "${DIST_DIR}/index.html" AND NOT "${VERSION}" STREQUAL "")
file(READ "${STAMP_FILE}" stamped)
string(STRIP "${stamped}" stamped)
if("${stamped}" STREQUAL "${VERSION}")
if(stamped STREQUAL "${stamp_key}")
set(stamp_ok TRUE)
endif()
endif()
set(have_assets FALSE)
if(EXISTS "${DIST_DIR}/index.html")
set(have_assets TRUE)
endif()
if(stamp_ok AND have_assets)
message(STATUS "UI: HF stamp '${stamped}' matches version, skipping HF fetch")
if(stamp_ok)
message(STATUS "UI: HF stamp matches '${stamp_key}', skipping HF fetch")
set(provisioned TRUE)
else()
hf_download("${VERSION}" HF_OK HF_RESOLVED)
if(HF_OK)
file(WRITE "${STAMP_FILE}" "${HF_RESOLVED}")
message(STATUS "UI: HF download succeeded, stamp updated (${HF_RESOLVED})")
file(WRITE "${STAMP_FILE}" "${stamp_key}")
message(STATUS "UI: HF download succeeded, stamp updated (${stamp_key}, resolved: ${HF_RESOLVED})")
set(provisioned TRUE)
else()
message(STATUS "UI: HF download failed")
+43 -32
View File
@@ -8,40 +8,44 @@ llama_add_compile_flags()
file(GLOB LLAMA_MODELS_SOURCES "models/*.cpp")
set(LLAMA_CORE_SOURCES
llama.cpp
llama-adapter.cpp
llama-arch.cpp
llama-batch.cpp
llama-chat.cpp
llama-context.cpp
llama-cparams.cpp
llama-grammar.cpp
llama-graph.cpp
llama-hparams.cpp
llama-impl.cpp
llama-io.cpp
llama-kv-cache.cpp
llama-kv-cache-iswa.cpp
llama-kv-cache-dsa.cpp
llama-kv-cache-dsa-iswa.cpp
llama-kv-cache-msa.cpp
llama-kv-cache-dsv4.cpp
llama-memory.cpp
llama-memory-hybrid.cpp
llama-memory-hybrid-iswa.cpp
llama-memory-hybrid-idx.cpp
llama-memory-recurrent.cpp
llama-mmap.cpp
llama-model-loader.cpp
llama-model-saver.cpp
llama-model.cpp
llama-quant.cpp
llama-sampler.cpp
llama-vocab.cpp
unicode-data.cpp
unicode.cpp
)
add_library(llama
../include/llama.h
llama.cpp
llama-adapter.cpp
llama-arch.cpp
llama-batch.cpp
llama-chat.cpp
llama-context.cpp
llama-cparams.cpp
llama-grammar.cpp
llama-graph.cpp
llama-hparams.cpp
llama-impl.cpp
llama-io.cpp
llama-kv-cache.cpp
llama-kv-cache-iswa.cpp
llama-kv-cache-dsa.cpp
llama-kv-cache-dsa-iswa.cpp
llama-kv-cache-msa.cpp
llama-kv-cache-dsv4.cpp
llama-memory.cpp
llama-memory-hybrid.cpp
llama-memory-hybrid-iswa.cpp
llama-memory-hybrid-idx.cpp
llama-memory-recurrent.cpp
llama-mmap.cpp
llama-model-loader.cpp
llama-model-saver.cpp
llama-model.cpp
llama-quant.cpp
llama-sampler.cpp
llama-vocab.cpp
unicode-data.cpp
unicode.cpp
${LLAMA_CORE_SOURCES}
unicode.h
${LLAMA_MODELS_SOURCES}
)
@@ -50,13 +54,20 @@ set_target_properties(llama PROPERTIES
VERSION ${LLAMA_VERSION_BASE}
SOVERSION ${LLAMA_VERSION_MAJOR}
MACHO_CURRENT_VERSION 0 # keep macOS linker from seeing oversized version number
UNITY_BUILD ON
UNITY_BUILD_BATCH_SIZE 16
)
# exclude non-model sources from unity build
set_source_files_properties(${LLAMA_CORE_SOURCES} ../include/llama.h unicode.h
PROPERTIES SKIP_UNITY_BUILD_INCLUSION ON)
configure_file(llama-version.h.in ${CMAKE_CURRENT_BINARY_DIR}/llama-version.h @ONLY)
target_include_directories(llama PRIVATE . ${CMAKE_CURRENT_BINARY_DIR})
target_include_directories(llama PUBLIC ../include)
target_compile_features (llama PRIVATE cxx_std_17) # don't bump
target_precompile_headers (llama PRIVATE models/models.h)
target_link_libraries(llama PUBLIC ggml)
+2
View File
@@ -666,7 +666,9 @@ void llama_context::sched_reserve() {
// need to implement a more robust mechanism that tries a few different inputs and analyzes the results
ggml_cgraph * gf = nullptr;
switch (model.arch) {
case LLM_ARCH_KIMI_LINEAR:
case LLM_ARCH_MINIMAX_01:
// [TAG_RESERVE_DIAG_DECAY]
// the `inp_diag_decay` tensor size scales with `n_seq_tokens^2` which
// makes `n_seqs == 1` use more memory for the compute graph compared to `n_seqs > 1`
gf = graph_reserve(n_tokens, 1, n_outputs_pp, mctx.get(), model.hparams.no_alloc);
+3 -3
View File
@@ -2644,9 +2644,9 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params,
filter = [&](uint32_t il) { return il >= hparams.n_layer(); };
}
if ((arch == LLM_ARCH_STEP35 || arch == LLM_ARCH_HY_V3 || arch == LLM_ARCH_GLM_DSA ||
arch == LLM_ARCH_MIMO2 || arch == LLM_ARCH_DEEPSEEK32) &&
hparams.n_layer_nextn > 0) {
// don't filter when n_layer_nextn is repurposed for a router layer the trunk attends
// or when a model is entirely n_layer_nextn layers and has no trunk
if (hparams.n_layer_nextn > 0 && hparams.n_layer() > 0 && hparams.router_layer < 0) {
if (params.ctx_type == LLAMA_CONTEXT_TYPE_MTP) {
filter = [&](uint32_t il) { return il >= hparams.n_layer(); };
} else {
+10 -10
View File
@@ -82,7 +82,7 @@ std::unique_ptr<llm_graph_context> llama_model_gemma3n::build_arch_graph(const l
}
// get 2D slice view from a 3D tensor, the idx corresponds to the 3rd dim
static ggml_tensor * ggml_view_2d_slice(ggml_context * ctx0, ggml_tensor * x, int idx) {
static ggml_tensor * gemma3n_view_2d_slice(ggml_context * ctx0, ggml_tensor * x, int idx) {
GGML_ASSERT(idx < (int) x->ne[2]);
return ggml_view_2d(ctx0, x, x->ne[0], x->ne[1], ggml_row_size(x->type, x->ne[0]),
idx * x->ne[0] * x->ne[1] * ggml_element_size(x));
@@ -139,7 +139,7 @@ llama_model_gemma3n::graph::graph(const llama_model & model, const llm_graph_par
ggml_tensor * predictions = altup_predict(cur, il); // [n_embd, n_tokens, n_altup]
// predicted value will go through self-attention and laurel
ggml_tensor * active_prediction = ggml_view_2d_slice(ctx0, predictions, i_altup_act); // [n_embd, n_tokens]
ggml_tensor * active_prediction = gemma3n_view_2d_slice(ctx0, predictions, i_altup_act); // [n_embd, n_tokens]
cur = active_prediction;
cb(cur, "active_prediction", il);
@@ -236,13 +236,13 @@ llama_model_gemma3n::graph::graph(const llama_model & model, const llm_graph_par
ggml_tensor * first_prediction; // [n_embd, n_tokens]
{
first_prediction = ggml_view_2d_slice(ctx0, corrected, i_altup_act); // [n_embd, n_tokens]
first_prediction = gemma3n_view_2d_slice(ctx0, corrected, i_altup_act); // [n_embd, n_tokens]
first_prediction = ggml_mul(ctx0, first_prediction, model.layers[il].altup_correct_scale);
first_prediction = build_lora_mm(model.layers[il].per_layer_inp_gate, first_prediction);
first_prediction = ggml_gelu(ctx0, first_prediction); // [n_embd_altup, n_tokens]
cb(first_prediction, "first_prediction_gated", il);
ggml_tensor * inp_this_layer = ggml_view_2d_slice(ctx0, inp_per_layer, il); // [n_embd_altup, n_tokens]
ggml_tensor * inp_this_layer = gemma3n_view_2d_slice(ctx0, inp_per_layer, il); // [n_embd_altup, n_tokens]
first_prediction = ggml_mul(ctx0, first_prediction, inp_this_layer); // [n_embd_altup, n_tokens]
cb(first_prediction, "first_prediction_scaled", il);
@@ -253,7 +253,7 @@ llama_model_gemma3n::graph::graph(const llama_model & model, const llm_graph_par
}
// equivalent to python code: corrected_predictions[1:] += first_prediction
{
ggml_tensor * slice_first = ggml_view_2d_slice(ctx0, corrected, 0);
ggml_tensor * slice_first = gemma3n_view_2d_slice(ctx0, corrected, 0);
ggml_tensor * slice_rest = ggml_view_3d(
ctx0, corrected, n_embd, n_tokens, n_altup - 1, ggml_row_size(corrected->type, n_embd),
ggml_row_size(corrected->type, n_embd * n_tokens), n_embd * n_tokens * ggml_element_size(corrected));
@@ -271,7 +271,7 @@ llama_model_gemma3n::graph::graph(const llama_model & model, const llm_graph_par
// cur now has multiple altup(s), we want to merge them back to 1 altup
{
ggml_tensor * target_magnitude = calc_magnitude(ggml_view_2d_slice(ctx0, cur, i_altup_act)); // [n_embd, n_tokens]
ggml_tensor * target_magnitude = calc_magnitude(gemma3n_view_2d_slice(ctx0, cur, i_altup_act)); // [n_embd, n_tokens]
// do a view to skip the first slice (active altup)
ggml_tensor * alt_slice =
ggml_view_3d(ctx0, cur, n_embd, n_tokens, n_altup - 1, ggml_row_size(cur->type, n_embd),
@@ -283,9 +283,9 @@ llama_model_gemma3n::graph::graph(const llama_model & model, const llm_graph_par
cb(altup_unembd, "altup_unembd", -1);
// equivalent to torch.mean(hidden_states, dim=0)
cur = ggml_view_2d_slice(ctx0, cur, 0); // [n_embd, n_tokens]
cur = gemma3n_view_2d_slice(ctx0, cur, 0); // [n_embd, n_tokens]
for (int i = 0; i < n_altup - 1; ++i) {
cur = ggml_add(ctx0, cur, ggml_view_2d_slice(ctx0, altup_unembd, i));
cur = ggml_add(ctx0, cur, gemma3n_view_2d_slice(ctx0, altup_unembd, i));
}
cur = ggml_scale(ctx0, cur, 1.0f / float(n_altup)); // [n_embd, n_tokens]
cb(cur, "unembd_merged", -1);
@@ -419,7 +419,7 @@ ggml_tensor * llama_model_gemma3n::graph::altup_compute_router_modalities(ggml_t
// input cur shape: [n_embd, n_tokens, n_altup]
// output shape: [n_embd, n_tokens, n_altup]
ggml_tensor * llama_model_gemma3n::graph::altup_predict(ggml_tensor * cur, int il) {
ggml_tensor * activated = ggml_view_2d_slice(ctx0, cur, i_altup_act); // [n_embd, n_tokens]
ggml_tensor * activated = gemma3n_view_2d_slice(ctx0, cur, i_altup_act); // [n_embd, n_tokens]
ggml_tensor * modalities = altup_compute_router_modalities(activated, il); // [n_altup, n_tokens]
cb(modalities, "modalities", il);
@@ -447,7 +447,7 @@ ggml_tensor * llama_model_gemma3n::graph::altup_correct(ggml_tensor * prediction
ggml_tensor * modalities = altup_compute_router_modalities(activated, il); // [n_altup, n_tokens]
cb(modalities, "modalities", il);
ggml_tensor * active_prediction = ggml_view_2d_slice(ctx0, predictions, i_altup_act);
ggml_tensor * active_prediction = gemma3n_view_2d_slice(ctx0, predictions, i_altup_act);
ggml_tensor * innovation = ggml_sub(ctx0, activated, active_prediction); // [n_embd, n_tokens]
cb(innovation, "innovation", il);
+2 -2
View File
@@ -145,7 +145,7 @@ std::unique_ptr<llm_graph_context> llama_model_gemma4::build_arch_graph(const ll
}
// get 2D slice view from a 3D tensor, the idx corresponds to the 3rd dim
static ggml_tensor * ggml_view_2d_slice(ggml_context * ctx0, ggml_tensor * x, int idx) {
static ggml_tensor * gemma4_view_2d_slice(ggml_context * ctx0, ggml_tensor * x, int idx) {
GGML_ASSERT(idx < (int) x->ne[2]);
return ggml_view_2d(ctx0, x, x->ne[0], x->ne[1], ggml_row_size(x->type, x->ne[0]),
idx * x->ne[0] * x->ne[1] * ggml_element_size(x));
@@ -372,7 +372,7 @@ llama_model_gemma4::graph::graph(const llama_model & model, const llm_graph_para
cur = build_lora_mm(model.layers[il].per_layer_inp_gate, cur); // [n_embd_per_layer, n_tokens]
cur = ggml_gelu(ctx0, cur);
ggml_tensor * inp_this_layer = ggml_view_2d_slice(ctx0, inp_per_layer, il); // [n_embd_per_layer, n_tokens]
ggml_tensor * inp_this_layer = gemma4_view_2d_slice(ctx0, inp_per_layer, il); // [n_embd_per_layer, n_tokens]
// TODO @ngxson : improve this
if (il == n_layer - 1 && inp_out_ids && cparams.embeddings_nextn_masked) {
+1
View File
@@ -229,6 +229,7 @@ llama_model_minimax_01::graph::graph(const llama_model & model, const llm_graph_
ggml_set_input(inp->inp_k_decay);
cb(inp->inp_k_decay, "k_decay_exp", -1);
// [TAG_RESERVE_DIAG_DECAY]
inp->inp_diag_decay = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_seq_tokens, n_seq_tokens, n_head, n_seqs);
ggml_set_input(inp->inp_diag_decay);
cb(inp->inp_diag_decay, "diag_decay_exp", -1);
+5 -5
View File
@@ -142,6 +142,11 @@ llama_model_plamo2::graph::graph(const llama_model & model, const llm_graph_para
cur = build_plamo2_attn_layer(inp_hybrid->get_attn(), inp_pos, cur, model, il);
}
if (il == n_layer - 1 && inp_out_ids) {
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
residual = ggml_get_rows(ctx0, residual, inp_out_ids);
}
// post_mixer_norm
cur = build_norm(cur, model.layers[il].attn_post_norm, NULL, LLM_NORM_RMS, il);
cb(cur, "attn_post_norm", il);
@@ -167,11 +172,6 @@ llama_model_plamo2::graph::graph(const llama_model & model, const llm_graph_para
cur = build_norm(cur, model.layers[il].ffn_post_norm, NULL, LLM_NORM_RMS, il);
cb(cur, "ffn_post_norm", il);
if (il == n_layer - 1 && inp_out_ids) {
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
residual = ggml_get_rows(ctx0, residual, inp_out_ids);
}
// residual connection
cur = ggml_add(ctx0, cur, residual);
cb(cur, "ffn_residual", il);
+1
View File
@@ -18,6 +18,7 @@ void llama_model_qwen3vl::load_arch_tensors(llama_model_loader &) {
int64_t n_vocab_out = n_vocab;
if (arch == LLM_ARCH_QWEN3TTS) {
// [TAG_LLAMA_N_VOCAB_OUT]
n_vocab_out = 3072;
}
+1
View File
@@ -1,6 +1,7 @@
*
!*.*
!snapshots/
!fusion/
*.o
ggml-common.h
**/*.swp
+8 -10
View File
@@ -10,7 +10,7 @@ function(llama_build source)
endif()
add_executable(${TEST_TARGET} ${TEST_SOURCES})
target_link_libraries(${TEST_TARGET} PRIVATE llama llama-common)
target_link_libraries(${TEST_TARGET} PRIVATE llama llama-common) # TODO: [TAG_TESTS_LLAMA_LINK]
if (LLAMA_TESTS_INSTALL)
install(TARGETS ${TEST_TARGET} RUNTIME)
endif()
@@ -196,7 +196,7 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS)
# llama_build_and_test(test-double-float.cpp) # SLOW
llama_build_and_test(test-llama-archs.cpp)
llama_build(test-llama-archs.cpp)
set(MODEL_DIR "${CMAKE_CURRENT_BINARY_DIR}/test-models/")
file(MAKE_DIRECTORY "${MODEL_DIR}")
@@ -255,6 +255,8 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS)
ARGS --models "${MODEL_DIR}"
)
set_tests_properties(test-save-load-state PROPERTIES FIXTURES_REQUIRED generate-models)
llama_build(test-fusion.cpp)
endif()
llama_build_and_test(test-chat-peg-parser.cpp peg-parser/simple-tokenize.cpp)
@@ -276,6 +278,8 @@ llama_build_and_test(
peg-parser/test-unicode.cpp
peg-parser/tests.h
)
target_precompile_headers(test-peg-parser PRIVATE peg-parser/tests.h)
if (NOT ${CMAKE_SYSTEM_PROCESSOR} MATCHES "s390x")
set(MODEL_NAME "tinyllamas/stories15M-q4_0.gguf")
@@ -306,15 +310,9 @@ if (NOT LLAMA_SANITIZE_ADDRESS AND NOT GGML_SCHED_NO_REALLOC)
# TODO: repair known memory leaks
llama_build_and_test(test-opt.cpp)
endif()
llama_build_and_test(test-backend-ops.cpp)
# the tensor API kernels come from a separate metallib - check they produce correct results
# ref: https://github.com/ggml-org/llama.cpp/issues/27473
if (GGML_METAL AND NOT GGML_METAL_EMBED_LIBRARY)
llama_test(test-backend-ops NAME test-backend-ops-metallib-tensor
ARGS test -b MTL0 -o MUL_MAT -p type_a=q6_K)
set_tests_properties(test-backend-ops-metallib-tensor PROPERTIES ENVIRONMENT GGML_METAL_TENSOR_ENABLE=1)
endif()
# TODO: make this test (and others) not link `libllama` as it is not needed [TAG_TESTS_LLAMA_LINK]
llama_build(test-backend-ops.cpp)
llama_build_and_test(test-model-load-cancel.cpp LABEL "model")
llama_build_and_test(test-autorelease.cpp LABEL "model")

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